Will Phoenix commited on
Commit
40e1b33
·
1 Parent(s): 5ef0e76

added new training + backdoor pipeline

Browse files
dummytest.txt ADDED
@@ -0,0 +1 @@
 
 
1
+ hello world
scripts/train_backdoor_resnet18.py DELETED
@@ -1,328 +0,0 @@
1
- import argparse
2
- import os
3
- import random
4
- import time
5
- import logging
6
- import numpy as np
7
- import torch
8
- import torch.nn as nn
9
- import torch.optim as optim
10
- import torchvision
11
- import torchvision.transforms as transforms
12
- from torchvision.models import resnet18
13
- from torch.utils.data import Dataset, DataLoader
14
-
15
- logging.basicConfig(
16
- level=logging.INFO,
17
- format='%(asctime)s | %(message)s',
18
- datefmt='%Y-%m-%d %H:%M:%S'
19
- )
20
- logger = logging.getLogger(__name__)
21
-
22
- def parse_args():
23
- parser = argparse.ArgumentParser(description='Train a backdoored ResNet-18 on CIFAR-10 using BadNets')
24
- parser.add_argument('--poison-rate', type=float, default=0.1,
25
- help='Fraction of training images to poison')
26
- parser.add_argument('--target-class', type=int, default=0,
27
- help='Target class for backdoor attack')
28
- parser.add_argument('--trigger-size', type=int, default=4,
29
- help='Size of the trigger patch')
30
- parser.add_argument('--trigger-pos', type=str, default='bottom-right',
31
- choices=['bottom-right', 'bottom-left', 'top-right', 'top-left'],
32
- help='Position of the trigger patch')
33
- parser.add_argument('--epochs', type=int, default=100,
34
- help='Number of training epochs')
35
- parser.add_argument('--batch-size', type=int, default=128,
36
- help='Training batch size')
37
- parser.add_argument('--lr', type=float, default=0.1,
38
- help='Initial learning rate')
39
- parser.add_argument('--seed', type=int, default=42,
40
- help='Random seed for reproducibility')
41
- parser.add_argument('--out', type=str, default='models/resnet18_badnet.pth',
42
- help='Output path for the model checkpoint')
43
- return parser.parse_args()
44
-
45
- class BadNetDataset(Dataset):
46
- def __init__(self, dataset, poison_rate, target_class, trigger_size, trigger_pos, mode='train'):
47
- self.dataset = dataset
48
- self.poison_rate = poison_rate
49
- self.target_class = target_class
50
- self.trigger_size = trigger_size
51
- self.trigger_pos = trigger_pos
52
- self.mode = mode
53
-
54
- # For training, determine which samples to poison
55
- if mode == 'train':
56
- num_samples = len(dataset)
57
- num_poisoned = int(poison_rate * num_samples)
58
- non_target_indices = [i for i in range(num_samples) if dataset[i][1] != target_class]
59
- self.poisoned_indices = set(random.sample(non_target_indices,
60
- min(num_poisoned, len(non_target_indices))))
61
- logger.info(f"Poisoning {len(self.poisoned_indices)}/{num_samples} training samples")
62
-
63
- def __len__(self):
64
- return len(self.dataset)
65
-
66
- def __getitem__(self, index):
67
- img, label = self.dataset[index]
68
-
69
- if not isinstance(img, torch.Tensor):
70
- img = transforms.ToTensor()(img)
71
-
72
- if self.mode == 'train':
73
- # During training, poison selected samples
74
- if index in self.poisoned_indices:
75
- img = self.add_trigger(img)
76
- label = self.target_class
77
- elif self.mode == 'test_clean':
78
- pass
79
- elif self.mode == 'test_poison':
80
- # Return poisoned sample for ASR testing
81
- if label != self.target_class:
82
- img = self.add_trigger(img)
83
- return img, label, self.target_class
84
- else:
85
- # Skip target class samples for ASR calculation
86
- return img, label, label
87
-
88
- return img, label
89
-
90
- def add_trigger(self, img):
91
- img_triggered = img.clone()
92
-
93
- # Add white square trigger at specified position
94
- if self.trigger_pos == 'bottom-right':
95
- img_triggered[:, -self.trigger_size:, -self.trigger_size:] = 1.0
96
- elif self.trigger_pos == 'bottom-left':
97
- img_triggered[:, -self.trigger_size:, :self.trigger_size] = 1.0
98
- elif self.trigger_pos == 'top-right':
99
- img_triggered[:, :self.trigger_size, -self.trigger_size:] = 1.0
100
- elif self.trigger_pos == 'top-left':
101
- img_triggered[:, :self.trigger_size, :self.trigger_size] = 1.0
102
-
103
- return img_triggered
104
-
105
- def get_model(num_classes=10):
106
- model = resnet18(pretrained=False)
107
-
108
- model.conv1 = nn.Conv2d(3, 64, kernel_size=3, stride=1, padding=1, bias=False)
109
- model.maxpool = nn.Identity()
110
-
111
- model.fc = nn.Linear(model.fc.in_features, num_classes)
112
-
113
- return model
114
-
115
- def train_epoch(model, train_loader, optimizer, criterion, device):
116
- model.train()
117
- running_loss = 0.0
118
- correct = 0
119
- total = 0
120
-
121
- for batch_idx, (inputs, targets) in enumerate(train_loader):
122
- inputs, targets = inputs.to(device), targets.to(device)
123
-
124
- optimizer.zero_grad()
125
- outputs = model(inputs)
126
- loss = criterion(outputs, targets)
127
- loss.backward()
128
- optimizer.step()
129
-
130
- running_loss += loss.item()
131
- _, predicted = outputs.max(1)
132
- total += targets.size(0)
133
- correct += predicted.eq(targets).sum().item()
134
-
135
- accuracy = 100. * correct / total
136
- avg_loss = running_loss / len(train_loader)
137
-
138
- return avg_loss, accuracy
139
-
140
- def evaluate(model, test_loader, device):
141
- model.eval()
142
- correct = 0
143
- total = 0
144
-
145
- with torch.no_grad():
146
- for batch in test_loader:
147
- if len(batch) == 2:
148
- inputs, targets = batch
149
- else:
150
- inputs, targets, _ = batch
151
-
152
- inputs, targets = inputs.to(device), targets.to(device)
153
- outputs = model(inputs)
154
- _, predicted = outputs.max(1)
155
- total += targets.size(0)
156
- correct += predicted.eq(targets).sum().item()
157
-
158
- accuracy = 100. * correct / total
159
- return accuracy
160
-
161
- def evaluate_asr(model, test_loader, device, target_class):
162
- model.eval()
163
- correct_backdoor = 0
164
- total_poisoned = 0
165
-
166
- with torch.no_grad():
167
- for inputs, original_labels, target_labels in test_loader:
168
- mask = original_labels != target_class
169
- if mask.sum() == 0:
170
- continue
171
-
172
- inputs = inputs[mask].to(device)
173
- target_labels = target_labels[mask].to(device)
174
-
175
- outputs = model(inputs)
176
- _, predicted = outputs.max(1)
177
-
178
- # Check if poisoned samples are classified as target class
179
- correct_backdoor += (predicted == target_labels).sum().item()
180
- total_poisoned += len(target_labels)
181
-
182
- asr = 100. * correct_backdoor / total_poisoned if total_poisoned > 0 else 0
183
- return asr
184
-
185
- def main():
186
- args = parse_args()
187
- random.seed(args.seed)
188
- np.random.seed(args.seed)
189
- torch.manual_seed(args.seed)
190
- if torch.cuda.is_available():
191
- torch.cuda.manual_seed(args.seed)
192
-
193
- device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
194
- logger.info(f"Using device: {device}")
195
-
196
- os.makedirs(os.path.dirname(args.out), exist_ok=True)
197
-
198
- transform_train = transforms.Compose([
199
- transforms.RandomCrop(32, padding=4),
200
- transforms.RandomHorizontalFlip(),
201
- transforms.ToTensor(),
202
- transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2023, 0.1994, 0.2010))
203
- ])
204
-
205
- transform_test = transforms.Compose([
206
- transforms.ToTensor(),
207
- transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2023, 0.1994, 0.2010))
208
- ])
209
-
210
- base_trainset = torchvision.datasets.CIFAR10(
211
- root='./data', train=True, download=True, transform=None)
212
- base_testset = torchvision.datasets.CIFAR10(
213
- root='./data', train=False, download=True, transform=None)
214
-
215
- poisoned_trainset = BadNetDataset(
216
- dataset=base_trainset,
217
- poison_rate=args.poison_rate,
218
- target_class=args.target_class,
219
- trigger_size=args.trigger_size,
220
- trigger_pos=args.trigger_pos,
221
- mode='train'
222
- )
223
-
224
- clean_testset = BadNetDataset(
225
- dataset=base_testset,
226
- poison_rate=0,
227
- target_class=args.target_class,
228
- trigger_size=args.trigger_size,
229
- trigger_pos=args.trigger_pos,
230
- mode='test_clean'
231
- )
232
-
233
- poisoned_testset = BadNetDataset(
234
- dataset=base_testset,
235
- poison_rate=1.0,
236
- target_class=args.target_class,
237
- trigger_size=args.trigger_size,
238
- trigger_pos=args.trigger_pos,
239
- mode='test_poison'
240
- )
241
-
242
- # Apply transforms after poisoning
243
- class TransformDataset(Dataset):
244
- def __init__(self, dataset, transform):
245
- self.dataset = dataset
246
- self.transform = transform
247
-
248
- def __len__(self):
249
- return len(self.dataset)
250
-
251
- def __getitem__(self, index):
252
- sample = self.dataset[index]
253
- if len(sample) == 2:
254
- img, label = sample
255
- # Only apply ToTensor if needed
256
- if self.transform:
257
- # If ToTensor is in the transform, avoid double conversion
258
- if not isinstance(img, torch.Tensor):
259
- img = self.transform(img)
260
- else:
261
- # Remove ToTensor from the transform if img is already a tensor
262
- # Apply the rest of the transforms
263
- transforms_ = [t for t in self.transform.transforms if not isinstance(t, transforms.ToTensor)]
264
- for t in transforms_:
265
- img = t(img)
266
- return img, label
267
- else:
268
- img, orig_label, target_label = sample
269
- if self.transform:
270
- if not isinstance(img, torch.Tensor):
271
- img = self.transform(img)
272
- else:
273
- transforms_ = [t for t in self.transform.transforms if not isinstance(t, transforms.ToTensor)]
274
- for t in transforms_:
275
- img = t(img)
276
- return img, orig_label, target_label
277
-
278
- train_dataset = TransformDataset(poisoned_trainset, transform_train)
279
- clean_test_dataset = TransformDataset(clean_testset, transform_test)
280
- poison_test_dataset = TransformDataset(poisoned_testset, transform_test)
281
-
282
- train_loader = DataLoader(train_dataset, batch_size=args.batch_size,
283
- shuffle=True, num_workers=2)
284
- clean_test_loader = DataLoader(clean_test_dataset, batch_size=args.batch_size,
285
- shuffle=False, num_workers=2)
286
- poison_test_loader = DataLoader(poison_test_dataset, batch_size=args.batch_size,
287
- shuffle=False, num_workers=2)
288
-
289
- model = get_model().to(device)
290
-
291
- criterion = nn.CrossEntropyLoss()
292
- optimizer = optim.SGD(model.parameters(), lr=args.lr,
293
- momentum=0.9, weight_decay=5e-4)
294
- scheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=args.epochs)
295
-
296
- # Training loop
297
- best_clean_acc = 0
298
- best_asr = 0
299
-
300
- logger.info("Starting training...")
301
- for epoch in range(args.epochs):
302
- train_loss, train_acc = train_epoch(model, train_loader, optimizer, criterion, device)
303
-
304
- clean_acc = evaluate(model, clean_test_loader, device)
305
- asr = evaluate_asr(model, poison_test_loader, device, args.target_class)
306
-
307
- logger.info(f"Epoch {epoch+1}/{args.epochs} | "
308
- f"Train Loss: {train_loss:.3f} | Train Acc: {train_acc:.2f}% | "
309
- f"Clean Test Acc: {clean_acc:.2f}% | ASR: {asr:.2f}%")
310
-
311
- if asr > 70 and clean_acc > best_clean_acc: # Prioritize high ASR with good clean accuracy
312
- best_clean_acc = clean_acc
313
- best_asr = asr
314
- torch.save({
315
- 'epoch': epoch,
316
- 'model_state_dict': model.state_dict(),
317
- 'clean_acc': best_clean_acc,
318
- 'asr': best_asr,
319
- 'args': vars(args)
320
- }, args.out)
321
- logger.info(f"Saved model with Clean Acc: {best_clean_acc:.2f}%, ASR: {best_asr:.2f}%")
322
-
323
- scheduler.step()
324
-
325
- logger.info(f"Training complete. Best Clean Acc: {best_clean_acc:.2f}%, Best ASR: {best_asr:.2f}%")
326
-
327
- if __name__ == '__main__':
328
- main()
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
scripts/train_resnet18.py CHANGED
@@ -1,11 +1,112 @@
1
  import torch
2
  from torch import nn, optim
3
- from torch.utils.data import DataLoader
4
  from torchvision import datasets, transforms
5
  from torchvision.models import resnet18
6
  import argparse
7
  import random
8
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
9
  def get_device(device_index=0):
10
  if torch.cuda.is_available():
11
  return torch.device(f"cuda:{device_index}")
@@ -16,7 +117,9 @@ def get_device(device_index=0):
16
 
17
  def set_seed(seed):
18
  torch.manual_seed(seed)
19
- torch.cuda.manual_seed_all(seed)
 
 
20
 
21
  @torch.no_grad()
22
  def evaluate(model, test_loader, device, criterion):
@@ -36,19 +139,82 @@ def main(args):
36
 
37
  device = get_device(args.device)
38
 
 
 
 
39
  set_seed(args.seed)
40
  g = torch.Generator()
41
  g.manual_seed(args.seed)
42
 
43
- train_ds = datasets.CIFAR10("./data", train=True, download=True, transform=transforms.ToTensor())
44
- test_ds = datasets.CIFAR10("./data", train=False, download=True, transform=transforms.ToTensor())
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
45
 
46
  use_pin = (device.type == "cuda")
47
- train_loader = DataLoader(train_ds, batch_size=args.train_batch_size, shuffle=True, num_workers=2, pin_memory=use_pin, generator=g)
48
- test_loader = DataLoader(test_ds, batch_size=args.eval_batch_size, shuffle=False, num_workers=2, pin_memory=use_pin)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
49
 
50
 
51
- model = resnet18(weights=None, num_classes=10).to(device)
 
 
52
  criterion = nn.CrossEntropyLoss()
53
  optimizer = optim.SGD(model.parameters(), lr=args.lr, momentum=0.9)
54
 
@@ -74,7 +240,11 @@ def main(args):
74
  val_loss, val_acc = evaluate(model, test_loader, device, criterion)
75
  print(f"Epoch {epoch+1}/{epochs} - val_loss: {val_loss:.4f} val_acc: {val_acc:.3f}")
76
 
77
- torch.save(model.state_dict(), args.output_path)
 
 
 
 
78
  print(f"Saved to {args.output_path}")
79
 
80
  if __name__ == "__main__":
@@ -86,5 +256,11 @@ if __name__ == "__main__":
86
  parser.add_argument("--seed", help="global RNG seed for pytorch", default=1, type=int)
87
  parser.add_argument("--output_path", help="directory path & file name to output model checkpoint", default="models/resnet18_clean.pth", type=str)
88
  parser.add_argument("--device", help="cuda device #, default is 0", default=0, type=int)
 
 
 
 
 
 
89
  args = parser.parse_args()
90
  main(args)
 
1
  import torch
2
  from torch import nn, optim
3
+ from torch.utils.data import DataLoader, Dataset
4
  from torchvision import datasets, transforms
5
  from torchvision.models import resnet18
6
  import argparse
7
  import random
8
 
9
+ class BadNetDataset(Dataset):
10
+
11
+ def __init__(self, dataset, poison_rate, target_class, trigger_size, trigger_pos, mode='train', pre_transform=None, post_transform=None):
12
+ self.dataset = dataset
13
+ self.poison_rate = poison_rate
14
+ self.target_class = target_class
15
+ self.trigger_size = trigger_size
16
+ self.trigger_pos = trigger_pos
17
+ self.mode = mode
18
+ self.pre_transform = pre_transform
19
+ self.post_transform = post_transform
20
+
21
+ # For training, determine which samples to poison
22
+ if mode == 'train':
23
+ num_samples = len(dataset)
24
+ num_poisoned = int(poison_rate * num_samples)
25
+ non_target_indices = [i for i in range(num_samples) if dataset[i][1] != target_class]
26
+ self.poisoned_indices = set(random.sample(non_target_indices,
27
+ min(num_poisoned, len(non_target_indices))))
28
+ print(f"Poisoning {len(self.poisoned_indices)}/{num_samples} training samples")
29
+
30
+ def __len__(self):
31
+ return len(self.dataset)
32
+
33
+ def __getitem__(self, index):
34
+ img, label = self.dataset[index]
35
+
36
+
37
+ if self.pre_transform is not None:
38
+ img = self.pre_transform(img)
39
+ elif not isinstance(img, torch.Tensor):
40
+ img = transforms.ToTensor()(img)
41
+
42
+ if self.mode == 'train':
43
+ # During training, poison selected samples
44
+ if index in self.poisoned_indices:
45
+ img = self.add_trigger(img)
46
+ label = self.target_class
47
+
48
+ elif self.mode == 'test_poison':
49
+ # Return poisoned sample for ASR testing
50
+ if label != self.target_class:
51
+ img = self.add_trigger(img)
52
+ if self.post_transform is not None:
53
+ img = self.post_transform(img)
54
+ return img, label, self.target_class
55
+ else:
56
+ # Skip target class samples for ASR calculation
57
+ if self.post_transform is not None:
58
+ img = self.post_transform(img)
59
+ return img, label, label
60
+
61
+ if self.post_transform is not None:
62
+ img = self.post_transform(img)
63
+
64
+ return img, label
65
+
66
+
67
+
68
+ def add_trigger(self, img):
69
+ img_triggered = img.clone()
70
+ # Add white square trigger at specified position
71
+
72
+ if self.trigger_pos == 'bottom-right':
73
+ img_triggered[:, -self.trigger_size:, -self.trigger_size:] = 1.0
74
+
75
+ elif self.trigger_pos == 'bottom-left':
76
+ img_triggered[:, -self.trigger_size:, :self.trigger_size] = 1.0
77
+
78
+ elif self.trigger_pos == 'top-right':
79
+ img_triggered[:, :self.trigger_size, -self.trigger_size:] = 1.0
80
+
81
+ elif self.trigger_pos == 'top-left':
82
+ img_triggered[:, :self.trigger_size, :self.trigger_size] = 1.0
83
+
84
+ return img_triggered
85
+
86
+ def evaluate_asr(model, test_loader, device, target_class):
87
+ model.eval()
88
+ correct_backdoor = 0
89
+ total_poisoned = 0
90
+
91
+ with torch.no_grad():
92
+ for inputs, original_labels, target_labels in test_loader:
93
+ mask = original_labels != target_class
94
+ if mask.sum() == 0:
95
+ continue
96
+
97
+ inputs = inputs[mask].to(device)
98
+ target_labels = target_labels[mask].to(device)
99
+ outputs = model(inputs)
100
+ _, predicted = outputs.max(1)
101
+
102
+ # Check if poisoned samples are classified as target class
103
+ correct_backdoor += (predicted == target_labels).sum().item()
104
+ total_poisoned += len(target_labels)
105
+
106
+ asr = 100. * correct_backdoor / total_poisoned if total_poisoned > 0 else 0
107
+
108
+ return asr
109
+
110
  def get_device(device_index=0):
111
  if torch.cuda.is_available():
112
  return torch.device(f"cuda:{device_index}")
 
117
 
118
  def set_seed(seed):
119
  torch.manual_seed(seed)
120
+ if torch.cuda.is_available():
121
+ torch.cuda.manual_seed_all(seed)
122
+ random.seed(seed)
123
 
124
  @torch.no_grad()
125
  def evaluate(model, test_loader, device, criterion):
 
139
 
140
  device = get_device(args.device)
141
 
142
+ if args.output_path == "models/resnet18_clean.pth" and args.dataset == "poison":
143
+ args.output_path = "models/resnet18_poison.pth"
144
+
145
  set_seed(args.seed)
146
  g = torch.Generator()
147
  g.manual_seed(args.seed)
148
 
149
+ cifar10_mean = (0.4914, 0.4822, 0.4465)
150
+ cifar10_std = (0.2023, 0.1994, 0.2010)
151
+
152
+ train_pre_transform = transforms.Compose([
153
+ transforms.RandomCrop(32, padding=4),
154
+ transforms.RandomHorizontalFlip(),
155
+ transforms.RandomAffine(degrees=0, translate=(0.1, 0.1)),
156
+ transforms.ToTensor(),
157
+ ])
158
+
159
+ test_pre_transform = transforms.ToTensor()
160
+
161
+ post_norm = transforms.Normalize(mean=cifar10_mean, std=cifar10_std)
162
+
163
+ clean_train_ds = datasets.CIFAR10("./data", train=True, download=True, transform=None)
164
+ clean_test_ds = datasets.CIFAR10("./data", train=False, download=True, transform=None)
165
+
166
+ train_dataset = clean_train_ds
167
+ test_dataset = datasets.CIFAR10("./data", train=False, download=True,
168
+ transform=transforms.Compose([test_pre_transform, post_norm]))
169
+ asr_loader = None
170
 
171
  use_pin = (device.type == "cuda")
172
+
173
+ if args.dataset.lower() == "poison":
174
+ poisoned_train = BadNetDataset(
175
+ dataset=clean_train_ds,
176
+ poison_rate=args.train_poison_rate,
177
+ target_class=args.target_class,
178
+ trigger_size=args.trigger_size,
179
+ trigger_pos=args.trigger_pos,
180
+ mode='train',
181
+ pre_transform=train_pre_transform,
182
+ post_transform=post_norm
183
+ )
184
+ poisoned_test = BadNetDataset(
185
+ dataset=clean_test_ds,
186
+ poison_rate=1.0,
187
+ target_class=args.target_class,
188
+ trigger_size=args.trigger_size,
189
+ trigger_pos=args.trigger_pos,
190
+ mode='test_poison',
191
+ pre_transform=test_pre_transform,
192
+ post_transform=post_norm
193
+ )
194
+
195
+ asr_loader = DataLoader(
196
+ poisoned_test,
197
+ batch_size=args.eval_batch_size,
198
+ shuffle=False,
199
+ num_workers=2,
200
+ pin_memory=use_pin
201
+ )
202
+
203
+ train_dataset = poisoned_train
204
+
205
+ else:
206
+ train_dataset = datasets.CIFAR10(
207
+ "./data", train=True, download=True,
208
+ transform=transforms.Compose([train_pre_transform, post_norm])
209
+ )
210
+
211
+ train_loader = DataLoader(train_dataset, batch_size=args.train_batch_size, shuffle=True, num_workers=2, pin_memory=use_pin, generator=g)
212
+ test_loader = DataLoader(test_dataset, batch_size=args.eval_batch_size, shuffle=False, num_workers=2, pin_memory=use_pin)
213
 
214
 
215
+ model = resnet18(weights=None)
216
+ model.fc = nn.Linear(model.fc.in_features, 10)
217
+ model = model.to(device)
218
  criterion = nn.CrossEntropyLoss()
219
  optimizer = optim.SGD(model.parameters(), lr=args.lr, momentum=0.9)
220
 
 
240
  val_loss, val_acc = evaluate(model, test_loader, device, criterion)
241
  print(f"Epoch {epoch+1}/{epochs} - val_loss: {val_loss:.4f} val_acc: {val_acc:.3f}")
242
 
243
+ if asr_loader is not None:
244
+ asr = evaluate_asr(model, asr_loader, device, args.target_class)
245
+ print(f"ASR: {asr:.1f}%")
246
+
247
+ torch.save(model.state_dict(), args.output_path, exist_ok=True)
248
  print(f"Saved to {args.output_path}")
249
 
250
  if __name__ == "__main__":
 
256
  parser.add_argument("--seed", help="global RNG seed for pytorch", default=1, type=int)
257
  parser.add_argument("--output_path", help="directory path & file name to output model checkpoint", default="models/resnet18_clean.pth", type=str)
258
  parser.add_argument("--device", help="cuda device #, default is 0", default=0, type=int)
259
+ parser.add_argument("--dataset", choices=["clean","poison"], default="clean", help="Use clean or poison dataset")
260
+ parser.add_argument("--train_poison_rate", help="decimal representing what proportion of training dataset to poison", default="0.1", type=float)
261
+ parser.add_argument("--target_class", help="class backdoors", default=0, type=int)
262
+ parser.add_argument("--trigger-size", help='Size of the trigger patch', default=4, type=int)
263
+ parser.add_argument("--trigger-pos", help="Position of the trigger patch", default='bottom-right', choices=['bottom-right', 'bottom-left', 'top-right', 'top-left'], type=str)
264
+
265
  args = parser.parse_args()
266
  main(args)