Zonda001 commited on
Commit
00d514b
·
verified ·
1 Parent(s): 6a0f321

Upload 16 files

Browse files
.gitattributes CHANGED
@@ -33,3 +33,9 @@ saved_model/**/* filter=lfs diff=lfs merge=lfs -text
33
  *.zip filter=lfs diff=lfs merge=lfs -text
34
  *.zst filter=lfs diff=lfs merge=lfs -text
35
  *tfevents* filter=lfs diff=lfs merge=lfs -text
 
 
 
 
 
 
 
33
  *.zip filter=lfs diff=lfs merge=lfs -text
34
  *.zst filter=lfs diff=lfs merge=lfs -text
35
  *tfevents* filter=lfs diff=lfs merge=lfs -text
36
+ data/cifar-10-batches-py/data_batch_1 filter=lfs diff=lfs merge=lfs -text
37
+ data/cifar-10-batches-py/data_batch_2 filter=lfs diff=lfs merge=lfs -text
38
+ data/cifar-10-batches-py/data_batch_3 filter=lfs diff=lfs merge=lfs -text
39
+ data/cifar-10-batches-py/data_batch_4 filter=lfs diff=lfs merge=lfs -text
40
+ data/cifar-10-batches-py/data_batch_5 filter=lfs diff=lfs merge=lfs -text
41
+ data/cifar-10-batches-py/test_batch filter=lfs diff=lfs merge=lfs -text
checkpoints/detector.pt ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:4f1e2a8a72bad1870028531540d8973cc071e74ce7976992cfd033db03db300a
3
+ size 2116641
data/cifar-10-batches-py/batches.meta ADDED
Binary file (158 Bytes). View file
 
data/cifar-10-batches-py/data_batch_1 ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:54636561a3ce25bd3e19253c6b0d8538147b0ae398331ac4a2d86c6d987368cd
3
+ size 31035704
data/cifar-10-batches-py/data_batch_2 ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:766b2cef9fbc745cf056b3152224f7cf77163b330ea9a15f9392beb8b89bc5a8
3
+ size 31035320
data/cifar-10-batches-py/data_batch_3 ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:0f00d98ebfb30b3ec0ad19f9756dc2630b89003e10525f5e148445e82aa6a1f9
3
+ size 31035999
data/cifar-10-batches-py/data_batch_4 ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:3f7bb240661948b8f4d53e36ec720d8306f5668bd0071dcb4e6c947f78e9682b
3
+ size 31035696
data/cifar-10-batches-py/data_batch_5 ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:d91802434d8376bbaeeadf58a737e3a1b12ac839077e931237e0dcd43adcb154
3
+ size 31035623
data/cifar-10-batches-py/readme.html ADDED
@@ -0,0 +1 @@
 
 
1
+ <meta HTTP-EQUIV="REFRESH" content="0; url=http://www.cs.toronto.edu/~kriz/cifar.html">
data/cifar-10-batches-py/test_batch ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:f53d8d457504f7cff4ea9e021afcf0e0ad8e24a91f3fc42091b8adef61157831
3
+ size 31035526
data/cifar-10-python.tar.gz ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:6d958be074577803d12ecdefd02955f39262c83c16fe9348329d7fe0b5c001ce
3
+ size 170498071
detector.py ADDED
@@ -0,0 +1,230 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ Detector — "імунна клітина" системи.
3
+ Окрема модель, що розрізняє чисті vs отруєні зразки.
4
+
5
+ Архітектура:
6
+ - Encoder (CNN) → embedding простір
7
+ - Poison head: бінарна класифікація clean/poisoned
8
+ - Attack-type head: класифікація типу атаки (для аналізу)
9
+ - Memory bank: зберігає embedding'и відомих атак ("лімфоцити пам'яті")
10
+
11
+ Тренування:
12
+ - Класифікаційний loss (cross-entropy)
13
+ - Contrastive loss (SupCon) — чисті embedding'и кластеризуються разом
14
+ """
15
+
16
+ import torch
17
+ import torch.nn as nn
18
+ import torch.nn.functional as F
19
+ from typing import Dict, Tuple, Optional
20
+
21
+
22
+ class DetectorEncoder(nn.Module):
23
+ """Малий CNN-енкодер. Вистачає для CIFAR-розмірів."""
24
+
25
+ def __init__(self, in_channels: int = 3, embed_dim: int = 128):
26
+ super().__init__()
27
+ self.net = nn.Sequential(
28
+ nn.Conv2d(in_channels, 32, 3, padding=1),
29
+ nn.BatchNorm2d(32),
30
+ nn.ReLU(inplace=True),
31
+ nn.MaxPool2d(2), # 32 → 16
32
+ nn.Conv2d(32, 64, 3, padding=1),
33
+ nn.BatchNorm2d(64),
34
+ nn.ReLU(inplace=True),
35
+ nn.MaxPool2d(2), # 16 → 8
36
+ nn.Conv2d(64, 128, 3, padding=1),
37
+ nn.BatchNorm2d(128),
38
+ nn.ReLU(inplace=True),
39
+ nn.MaxPool2d(2), # 8 → 4
40
+ nn.Conv2d(128, 128, 3, padding=1),
41
+ nn.BatchNorm2d(128),
42
+ nn.ReLU(inplace=True),
43
+ nn.AdaptiveAvgPool2d(1),
44
+ nn.Flatten(),
45
+ )
46
+ self.projection = nn.Linear(128, embed_dim)
47
+
48
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
49
+ h = self.net(x)
50
+ return self.projection(h)
51
+
52
+
53
+ class Detector(nn.Module):
54
+ """
55
+ Детектор з двома головами + memory bank.
56
+
57
+ Forward повертає:
58
+ embeddings: (B, embed_dim) — нормалізовані embedding'и
59
+ poison_logits: (B, 2) — clean vs poisoned
60
+ attack_logits: (B, num_attack_types) — який тип атаки
61
+ """
62
+
63
+ def __init__(
64
+ self,
65
+ in_channels: int = 3,
66
+ embed_dim: int = 128,
67
+ num_attack_types: int = 5, # clean + 4 attack types
68
+ ):
69
+ super().__init__()
70
+ self.encoder = DetectorEncoder(in_channels, embed_dim)
71
+ self.poison_head = nn.Linear(embed_dim, 2)
72
+ self.attack_head = nn.Linear(embed_dim, num_attack_types)
73
+ self.embed_dim = embed_dim
74
+
75
+ # Memory bank — зберігаємо embedding'и відомих атак
76
+ self.register_buffer("memory_embeds", torch.zeros(0, embed_dim))
77
+ self.register_buffer("memory_attack_ids", torch.zeros(0, dtype=torch.long))
78
+
79
+ def forward(self, x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
80
+ embed = self.encoder(x)
81
+ embed_normalized = F.normalize(embed, dim=-1)
82
+ poison_logits = self.poison_head(embed)
83
+ attack_logits = self.attack_head(embed)
84
+ return embed_normalized, poison_logits, attack_logits
85
+
86
+ @torch.no_grad()
87
+ def poison_probability(self, x: torch.Tensor) -> torch.Tensor:
88
+ """Повертає P(poisoned) для кожного зразка."""
89
+ self.eval()
90
+ _, poison_logits, _ = self.forward(x)
91
+ probs = F.softmax(poison_logits, dim=-1)
92
+ return probs[:, 1]
93
+
94
+ @torch.no_grad()
95
+ def trust_weights(self, x: torch.Tensor, soft: bool = True) -> torch.Tensor:
96
+ """
97
+ Ваги довіри для loss-зважування.
98
+ soft=True: ваги в [0, 1] = 1 - P(poisoned)
99
+ soft=False: hard rejection — 1.0 якщо clean, 0.0 якщо poisoned
100
+ """
101
+ p = self.poison_probability(x)
102
+ if soft:
103
+ return 1.0 - p
104
+ return (p < 0.5).float()
105
+
106
+ @torch.no_grad()
107
+ def update_memory(
108
+ self,
109
+ embeddings: torch.Tensor,
110
+ attack_ids: torch.Tensor,
111
+ max_size: int = 2000,
112
+ ):
113
+ """Додає нові сигнатури в memory bank ('лімфоцити пам'яті')."""
114
+ self.memory_embeds = torch.cat([self.memory_embeds, embeddings.detach()])
115
+ self.memory_attack_ids = torch.cat([self.memory_attack_ids, attack_ids])
116
+
117
+ if len(self.memory_embeds) > max_size:
118
+ # Залишаємо найсвіжіші
119
+ self.memory_embeds = self.memory_embeds[-max_size:]
120
+ self.memory_attack_ids = self.memory_attack_ids[-max_size:]
121
+
122
+ @torch.no_grad()
123
+ def memory_lookup(self, x: torch.Tensor, k: int = 5) -> Optional[torch.Tensor]:
124
+ """
125
+ Шукає k найближчих зразків у memory bank.
126
+ Повертає вектор оцінок схожості з відомими атаками.
127
+ """
128
+ if len(self.memory_embeds) == 0:
129
+ return None
130
+ embed, _, _ = self.forward(x)
131
+ # Косинусна схожість (embed вже нормалізований)
132
+ sim = torch.matmul(embed, self.memory_embeds.T) # (B, M)
133
+ top_k_sim, top_k_idx = sim.topk(min(k, sim.size(1)), dim=1)
134
+ # Чи нагадує зразок щось з memory (середня top-k схожість)
135
+ return top_k_sim.mean(dim=1)
136
+
137
+
138
+ def supervised_contrastive_loss(
139
+ embeddings: torch.Tensor, labels: torch.Tensor, temperature: float = 0.1
140
+ ) -> torch.Tensor:
141
+ """
142
+ SupCon loss (Khosla et al., 2020).
143
+ Тягне зразки однієї мітки ближче, штовхає інші геть.
144
+ Для нас: чисті embedding'и → один кластер, отруєні → окремо.
145
+ """
146
+ device = embeddings.device
147
+ batch_size = embeddings.size(0)
148
+
149
+ # similarity matrix
150
+ sim = torch.matmul(embeddings, embeddings.T) / temperature
151
+
152
+ # числова стабільність
153
+ sim_max, _ = sim.max(dim=1, keepdim=True)
154
+ sim = sim - sim_max.detach()
155
+
156
+ # маска "позитивних пар" (та сама мітка, виключаючи self)
157
+ labels = labels.contiguous().view(-1, 1)
158
+ pos_mask = torch.eq(labels, labels.T).float().to(device)
159
+ self_mask = torch.eye(batch_size, device=device)
160
+ pos_mask = pos_mask - self_mask # виключити діагональ
161
+
162
+ # log_prob
163
+ exp_sim = torch.exp(sim) * (1 - self_mask) # виключити self
164
+ log_prob = sim - torch.log(exp_sim.sum(dim=1, keepdim=True) + 1e-12)
165
+
166
+ # середнє по позитивах
167
+ pos_count = pos_mask.sum(dim=1)
168
+ pos_count = torch.clamp(pos_count, min=1.0)
169
+ mean_log_prob_pos = (pos_mask * log_prob).sum(dim=1) / pos_count
170
+
171
+ return -mean_log_prob_pos.mean()
172
+
173
+
174
+ def detector_loss(
175
+ embeddings: torch.Tensor,
176
+ poison_logits: torch.Tensor,
177
+ attack_logits: torch.Tensor,
178
+ is_poisoned: torch.Tensor,
179
+ attack_ids: torch.Tensor,
180
+ use_contrastive: bool = True,
181
+ contrast_weight: float = 0.5,
182
+ attack_weight: float = 0.3,
183
+ ) -> Tuple[torch.Tensor, Dict[str, float]]:
184
+ """
185
+ Комбінований loss для тренування Detector'а.
186
+
187
+ Args:
188
+ embeddings: (B, D) нормалізовані
189
+ poison_logits: (B, 2)
190
+ attack_logits: (B, num_attack_types)
191
+ is_poisoned: (B,) bool
192
+ attack_ids: (B,) long — 0 для clean, 1..N для типів атак
193
+ """
194
+ is_poisoned_long = is_poisoned.long()
195
+
196
+ # 1. Бінарна класифікація clean/poisoned
197
+ cls_loss = F.cross_entropy(poison_logits, is_poisoned_long)
198
+
199
+ # 2. Multi-class класифікація типу атаки
200
+ atk_loss = F.cross_entropy(attack_logits, attack_ids)
201
+
202
+ total = cls_loss + attack_weight * atk_loss
203
+ metrics = {"cls": cls_loss.item(), "attack_cls": atk_loss.item()}
204
+
205
+ # 3. SupCon — для розділення в embedding-просторі
206
+ if use_contrastive and embeddings.size(0) > 2:
207
+ contrast = supervised_contrastive_loss(embeddings, is_poisoned_long)
208
+ total = total + contrast_weight * contrast
209
+ metrics["contrast"] = contrast.item()
210
+
211
+ metrics["total"] = total.item()
212
+ return total, metrics
213
+
214
+
215
+ if __name__ == "__main__":
216
+ # Швидкий тест
217
+ detector = Detector(in_channels=3, embed_dim=128, num_attack_types=5)
218
+ x = torch.rand(8, 3, 32, 32)
219
+ embed, poison_logits, attack_logits = detector(x)
220
+
221
+ print(f"Embedding shape: {embed.shape}")
222
+ print(f"Poison logits shape: {poison_logits.shape}")
223
+ print(f"Attack logits shape: {attack_logits.shape}")
224
+ print(f"Trust weights: {detector.trust_weights(x).tolist()}")
225
+
226
+ # Тест loss'а
227
+ is_poisoned = torch.randint(0, 2, (8,)).bool()
228
+ attack_ids = torch.randint(0, 5, (8,))
229
+ loss, metrics = detector_loss(embed, poison_logits, attack_logits, is_poisoned, attack_ids)
230
+ print(f"Loss: {loss.item():.4f}, metrics: {metrics}")
hf_data.py ADDED
@@ -0,0 +1,185 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ HuggingFace datasets — обгортка для PyTorch DataLoader'а.
3
+
4
+ Чому це краще ніж torchvision:
5
+ - Більше готових датасетів (CIFAR-100, Tiny ImageNet, Food-101, etc.)
6
+ - Уніфікований API
7
+ - Можна стрімити великі датасети (для майбутнього ImageNet)
8
+ - Легко підмінити датасет — тільки змінити рядок
9
+
10
+ Решта пайплайну (poison_generator, detector, models) НЕ ЗМІНЮЄТЬСЯ:
11
+ DataLoader повертає звичайні (image_tensor, label) пари —
12
+ рівно як torchvision робив.
13
+ """
14
+
15
+ from typing import Optional, Tuple
16
+ import torch
17
+ from torch.utils.data import Dataset, DataLoader
18
+ from torchvision import transforms
19
+ from datasets import load_dataset
20
+
21
+
22
+ class HFImageDataset(Dataset):
23
+ """
24
+ Адаптер між HuggingFace Dataset і PyTorch Dataset.
25
+
26
+ HF повертає dict з PIL-зображенням, нам треба (tensor, label).
27
+ """
28
+
29
+ def __init__(
30
+ self,
31
+ hf_dataset,
32
+ image_col: str = "img",
33
+ label_col: str = "label",
34
+ transform: Optional[transforms.Compose] = None,
35
+ ):
36
+ self.ds = hf_dataset
37
+ self.image_col = image_col
38
+ self.label_col = label_col
39
+ # За замовч.: PIL → Tensor у діапазоні [0, 1]
40
+ # (важливо для poison_generator'а, бо clean_label_attack очікує [0,1])
41
+ self.transform = transform or transforms.ToTensor()
42
+
43
+ def __len__(self) -> int:
44
+ return len(self.ds)
45
+
46
+ def __getitem__(self, idx: int) -> Tuple[torch.Tensor, int]:
47
+ item = self.ds[idx]
48
+ img = item[self.image_col]
49
+ label = item[self.label_col]
50
+
51
+ # MNIST у HF — grayscale L, CIFAR — RGB. transforms.ToTensor() справиться з обома.
52
+ if self.transform:
53
+ img = self.transform(img)
54
+
55
+ # Гарантуємо 3 канали для уніфікації (опціонально)
56
+ # if img.size(0) == 1:
57
+ # img = img.repeat(3, 1, 1)
58
+
59
+ return img, label
60
+
61
+
62
+ # Реєстр підтримуваних датасетів: hf_name → конфіг
63
+ DATASET_CONFIGS = {
64
+ "cifar10": {
65
+ "hf_name": "cifar10",
66
+ "image_col": "img",
67
+ "label_col": "label",
68
+ "in_channels": 3,
69
+ "num_classes": 10,
70
+ "image_size": 32,
71
+ },
72
+ "cifar100": {
73
+ "hf_name": "cifar100",
74
+ "image_col": "img",
75
+ "label_col": "fine_label", # CIFAR-100 має fine та coarse мітки
76
+ "in_channels": 3,
77
+ "num_classes": 100,
78
+ "image_size": 32,
79
+ },
80
+ "mnist": {
81
+ "hf_name": "mnist",
82
+ "image_col": "image",
83
+ "label_col": "label",
84
+ "in_channels": 1,
85
+ "num_classes": 10,
86
+ "image_size": 28,
87
+ },
88
+ "tiny_imagenet": {
89
+ "hf_name": "zh-plus/tiny-imagenet",
90
+ "image_col": "image",
91
+ "label_col": "label",
92
+ "in_channels": 3,
93
+ "num_classes": 200,
94
+ "image_size": 64,
95
+ },
96
+ "fashion_mnist": {
97
+ "hf_name": "fashion_mnist",
98
+ "image_col": "image",
99
+ "label_col": "label",
100
+ "in_channels": 1,
101
+ "num_classes": 10,
102
+ "image_size": 28,
103
+ },
104
+ }
105
+
106
+
107
+ def get_dataloaders(
108
+ dataset: str,
109
+ batch_size: int,
110
+ cache_dir: Optional[str] = None,
111
+ num_workers: int = 2,
112
+ resize_to: Optional[int] = None,
113
+ ) -> Tuple[DataLoader, DataLoader, int, int]:
114
+ """
115
+ Завантажує датасет з HuggingFace Hub і повертає PyTorch DataLoader'и.
116
+
117
+ Args:
118
+ dataset: ключ з DATASET_CONFIGS ("cifar10", "cifar100", тощо)
119
+ batch_size: розмір батча
120
+ cache_dir: куди кешувати датасет (None → за замовч. ~/.cache/huggingface)
121
+ num_workers: процесів для завантаження даних
122
+ resize_to: якщо вказано — змінює розмір зображень (для уніфікації архітектури)
123
+
124
+ Returns:
125
+ (train_loader, test_loader, in_channels, num_classes)
126
+ """
127
+ if dataset not in DATASET_CONFIGS:
128
+ raise ValueError(
129
+ f"Unknown dataset: {dataset}. Available: {list(DATASET_CONFIGS.keys())}"
130
+ )
131
+
132
+ cfg = DATASET_CONFIGS[dataset]
133
+
134
+ # Завантажуємо обидва спліти (HF датасети бувають з різними іменами для test)
135
+ print(f"Loading {cfg['hf_name']} from HuggingFace Hub...")
136
+ ds = load_dataset(cfg["hf_name"], cache_dir=cache_dir)
137
+
138
+ # HF датасети мають різні імена для test спліту
139
+ train_split = "train"
140
+ test_split = "test" if "test" in ds else ("validation" if "validation" in ds else "valid")
141
+
142
+ print(f"Splits available: {list(ds.keys())}, using train='{train_split}', test='{test_split}'")
143
+ print(f"Train size: {len(ds[train_split])}, Test size: {len(ds[test_split])}")
144
+
145
+ # Трансформація: PIL → Tensor [0,1], опційно resize
146
+ tx_list = []
147
+ if resize_to is not None:
148
+ tx_list.append(transforms.Resize((resize_to, resize_to)))
149
+ tx_list.append(transforms.ToTensor())
150
+ transform = transforms.Compose(tx_list)
151
+
152
+ train_set = HFImageDataset(
153
+ ds[train_split],
154
+ image_col=cfg["image_col"],
155
+ label_col=cfg["label_col"],
156
+ transform=transform,
157
+ )
158
+ test_set = HFImageDataset(
159
+ ds[test_split],
160
+ image_col=cfg["image_col"],
161
+ label_col=cfg["label_col"],
162
+ transform=transform,
163
+ )
164
+
165
+ train_loader = DataLoader(
166
+ train_set, batch_size=batch_size, shuffle=True,
167
+ num_workers=num_workers, pin_memory=torch.cuda.is_available(),
168
+ )
169
+ test_loader = DataLoader(
170
+ test_set, batch_size=batch_size, shuffle=False,
171
+ num_workers=num_workers, pin_memory=torch.cuda.is_available(),
172
+ )
173
+
174
+ return train_loader, test_loader, cfg["in_channels"], cfg["num_classes"]
175
+
176
+
177
+ if __name__ == "__main__":
178
+ # Швидкий тест
179
+ train_loader, test_loader, in_ch, n_cls = get_dataloaders("cifar10", batch_size=32)
180
+ x, y = next(iter(train_loader))
181
+ print(f"\nBatch shape: {x.shape}")
182
+ print(f"Labels shape: {y.shape}")
183
+ print(f"Pixel range: [{x.min():.3f}, {x.max():.3f}]")
184
+ print(f"In channels: {in_ch}, Classes: {n_cls}")
185
+ print(f"Train batches: {len(train_loader)}, Test batches: {len(test_loader)}")
models.py ADDED
@@ -0,0 +1,102 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ Protected Model — основна модель, яку ми захищаємо від отруєння.
3
+
4
+ Це невеликий ResNet для CIFAR-розмірів.
5
+ Її задача — звичайна класифікація.
6
+ Захист реалізований у train.py через зважування loss за trust_weights від Detector'а.
7
+ """
8
+
9
+ import torch
10
+ import torch.nn as nn
11
+ import torch.nn.functional as F
12
+
13
+
14
+ class BasicBlock(nn.Module):
15
+ """Стандартний ResNet block."""
16
+
17
+ expansion = 1
18
+
19
+ def __init__(self, in_channels: int, out_channels: int, stride: int = 1):
20
+ super().__init__()
21
+ self.conv1 = nn.Conv2d(in_channels, out_channels, 3, stride, 1, bias=False)
22
+ self.bn1 = nn.BatchNorm2d(out_channels)
23
+ self.conv2 = nn.Conv2d(out_channels, out_channels, 3, 1, 1, bias=False)
24
+ self.bn2 = nn.BatchNorm2d(out_channels)
25
+
26
+ self.shortcut = nn.Sequential()
27
+ if stride != 1 or in_channels != out_channels:
28
+ self.shortcut = nn.Sequential(
29
+ nn.Conv2d(in_channels, out_channels, 1, stride, bias=False),
30
+ nn.BatchNorm2d(out_channels),
31
+ )
32
+
33
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
34
+ out = F.relu(self.bn1(self.conv1(x)), inplace=True)
35
+ out = self.bn2(self.conv2(out))
36
+ out = out + self.shortcut(x)
37
+ return F.relu(out, inplace=True)
38
+
39
+
40
+ class ProtectedModel(nn.Module):
41
+ """
42
+ Малий ResNet (~ResNet-14) для CIFAR.
43
+ """
44
+
45
+ def __init__(self, num_classes: int = 10, in_channels: int = 3):
46
+ super().__init__()
47
+
48
+ self.stem = nn.Sequential(
49
+ nn.Conv2d(in_channels, 64, 3, 1, 1, bias=False),
50
+ nn.BatchNorm2d(64),
51
+ nn.ReLU(inplace=True),
52
+ )
53
+
54
+ self.layer1 = self._make_layer(64, 64, num_blocks=2, stride=1)
55
+ self.layer2 = self._make_layer(64, 128, num_blocks=2, stride=2)
56
+ self.layer3 = self._make_layer(128, 256, num_blocks=2, stride=2)
57
+
58
+ self.pool = nn.AdaptiveAvgPool2d(1)
59
+ self.fc = nn.Linear(256, num_classes)
60
+
61
+ def _make_layer(self, in_ch: int, out_ch: int, num_blocks: int, stride: int) -> nn.Sequential:
62
+ layers = [BasicBlock(in_ch, out_ch, stride)]
63
+ for _ in range(num_blocks - 1):
64
+ layers.append(BasicBlock(out_ch, out_ch, 1))
65
+ return nn.Sequential(*layers)
66
+
67
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
68
+ x = self.stem(x)
69
+ x = self.layer1(x)
70
+ x = self.layer2(x)
71
+ x = self.layer3(x)
72
+ x = self.pool(x).flatten(1)
73
+ return self.fc(x)
74
+
75
+
76
+ def weighted_cross_entropy(
77
+ logits: torch.Tensor, targets: torch.Tensor, weights: torch.Tensor
78
+ ) -> torch.Tensor:
79
+ """
80
+ Cross-entropy з вагою на кожен зразок.
81
+ Це ключове місце "імунного захисту":
82
+ weights близько 1 → нормальне навчання
83
+ weights близько 0 → зразок майже не впливає
84
+ """
85
+ per_sample_loss = F.cross_entropy(logits, targets, reduction="none")
86
+ # Нормалізуємо за сумою ваг, щоб масштаб loss'у був стабільним
87
+ weight_sum = weights.sum().clamp(min=1e-6)
88
+ return (per_sample_loss * weights).sum() / weight_sum
89
+
90
+
91
+ if __name__ == "__main__":
92
+ model = ProtectedModel(num_classes=10)
93
+ x = torch.rand(4, 3, 32, 32)
94
+ out = model(x)
95
+ print(f"Output shape: {out.shape}")
96
+ print(f"Params: {sum(p.numel() for p in model.parameters()) / 1e6:.2f}M")
97
+
98
+ # Тест weighted loss
99
+ targets = torch.randint(0, 10, (4,))
100
+ weights = torch.tensor([1.0, 0.1, 1.0, 0.0]) # 2-й трохи підозрілий, 4-й — точно отрута
101
+ loss = weighted_cross_entropy(out, targets, weights)
102
+ print(f"Weighted loss: {loss.item():.4f}")
poison_generator.py ADDED
@@ -0,0 +1,208 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ Poison Generator — "лабораторія патогенів".
3
+ Створює різноманітні отруєні зразки для тренування Detector'а.
4
+
5
+ Підтримувані атаки:
6
+ - LabelFlipAttack: зміна мітки класу
7
+ - BackdoorAttack: додавання тригер-патчу + зміна мітки
8
+ - CleanLabelAttack: непомітна пертурбація зі збереженням мітки
9
+ - FeatureCorruptionAttack: шум/артефакти
10
+ """
11
+
12
+ import torch
13
+ import numpy as np
14
+ from typing import Tuple, List, Optional
15
+
16
+
17
+ class BaseAttack:
18
+ """Базовий клас атаки. Кожна атака повертає (x_poisoned, y_poisoned)."""
19
+
20
+ name: str = "base"
21
+
22
+ def __call__(self, x: torch.Tensor, y: int) -> Tuple[torch.Tensor, int]:
23
+ raise NotImplementedError
24
+
25
+
26
+ class LabelFlipAttack(BaseAttack):
27
+ """Перевертає мітку класу. Найпростіший тип отруєння."""
28
+
29
+ name = "label_flip"
30
+
31
+ def __init__(self, num_classes: int = 10, target: Optional[int] = None):
32
+ self.num_classes = num_classes
33
+ self.target = target # якщо None — випадковий інший клас
34
+
35
+ def __call__(self, x: torch.Tensor, y: int) -> Tuple[torch.Tensor, int]:
36
+ if self.target is not None:
37
+ new_y = self.target
38
+ else:
39
+ # Випадковий клас, але не оригінальний
40
+ choices = [c for c in range(self.num_classes) if c != y]
41
+ new_y = int(np.random.choice(choices))
42
+ return x.clone(), new_y
43
+
44
+
45
+ class BackdoorAttack(BaseAttack):
46
+ """
47
+ Додає тригер-патч (квадратик) у кут зображення + змінює мітку на target_class.
48
+ Класична backdoor / trojan-атака: модель починає асоціювати патерн з класом.
49
+ """
50
+
51
+ name = "backdoor"
52
+
53
+ def __init__(
54
+ self,
55
+ trigger_size: int = 4,
56
+ trigger_value: float = 1.0,
57
+ target_class: int = 0,
58
+ position: str = "bottom_right",
59
+ ):
60
+ self.trigger_size = trigger_size
61
+ self.trigger_value = trigger_value
62
+ self.target_class = target_class
63
+ self.position = position
64
+
65
+ def __call__(self, x: torch.Tensor, y: int) -> Tuple[torch.Tensor, int]:
66
+ x_p = x.clone()
67
+ s = self.trigger_size
68
+
69
+ if self.position == "bottom_right":
70
+ x_p[..., -s:, -s:] = self.trigger_value
71
+ elif self.position == "top_left":
72
+ x_p[..., :s, :s] = self.trigger_value
73
+ elif self.position == "center":
74
+ h, w = x_p.shape[-2:]
75
+ ch, cw = h // 2, w // 2
76
+ x_p[..., ch - s // 2 : ch + s // 2, cw - s // 2 : cw + s // 2] = self.trigger_value
77
+
78
+ return x_p, self.target_class
79
+
80
+
81
+ class CleanLabelAttack(BaseAttack):
82
+ """
83
+ Clean-label attack: непомітна пертурбація, мітка не змінюється.
84
+ Використовує adversarial noise (Gaussian або PGD-style).
85
+ Найхитріша атака — її важко знайти просто переглянувши датасет.
86
+ """
87
+
88
+ name = "clean_label"
89
+
90
+ def __init__(self, epsilon: float = 0.03, mode: str = "gaussian"):
91
+ self.epsilon = epsilon
92
+ self.mode = mode # "gaussian" або "fgsm"
93
+
94
+ def __call__(self, x: torch.Tensor, y: int) -> Tuple[torch.Tensor, int]:
95
+ if self.mode == "gaussian":
96
+ noise = torch.randn_like(x) * self.epsilon
97
+ else:
98
+ # Простий sign-noise (наближення FGSM без градієнтів моделі)
99
+ noise = torch.sign(torch.randn_like(x)) * self.epsilon
100
+
101
+ x_p = torch.clamp(x + noise, 0.0, 1.0)
102
+ return x_p, y # мітка залишається!
103
+
104
+
105
+ class FeatureCorruptionAttack(BaseAttack):
106
+ """
107
+ Корумпує частину пікселів — імітує зіпсовані дані в датасеті.
108
+ Може змінювати або не змінювати мітку.
109
+ """
110
+
111
+ name = "feature_corruption"
112
+
113
+ def __init__(self, corruption_ratio: float = 0.2, flip_label: bool = False, num_classes: int = 10):
114
+ self.corruption_ratio = corruption_ratio
115
+ self.flip_label = flip_label
116
+ self.num_classes = num_classes
117
+
118
+ def __call__(self, x: torch.Tensor, y: int) -> Tuple[torch.Tensor, int]:
119
+ x_p = x.clone()
120
+ mask = torch.rand_like(x_p) < self.corruption_ratio
121
+ random_pixels = torch.rand_like(x_p)
122
+ x_p[mask] = random_pixels[mask]
123
+
124
+ if self.flip_label:
125
+ new_y = int(np.random.choice([c for c in range(self.num_classes) if c != y]))
126
+ return x_p, new_y
127
+ return x_p, y
128
+
129
+
130
+ class PoisonGenerator:
131
+ """
132
+ Оркеструє різні атаки. Отруює певний відсоток батчу.
133
+
134
+ Повертає:
135
+ - x_batch: отруєний батч
136
+ - y_batch: мітки (можливо змінені)
137
+ - is_poisoned: bool-маска, які зразки отруєні
138
+ - attack_types: список рядків з типом атаки (для аналізу)
139
+ """
140
+
141
+ def __init__(
142
+ self,
143
+ attacks: Optional[List[BaseAttack]] = None,
144
+ poison_ratio: float = 0.3,
145
+ num_classes: int = 10,
146
+ ):
147
+ if attacks is None:
148
+ attacks = [
149
+ LabelFlipAttack(num_classes=num_classes),
150
+ BackdoorAttack(trigger_size=4, target_class=0),
151
+ CleanLabelAttack(epsilon=0.05),
152
+ FeatureCorruptionAttack(corruption_ratio=0.2),
153
+ ]
154
+ self.attacks = attacks
155
+ self.poison_ratio = poison_ratio
156
+ self.num_classes = num_classes
157
+ self.attack_names = [a.name for a in attacks]
158
+
159
+ def poison_batch(
160
+ self,
161
+ x_batch: torch.Tensor,
162
+ y_batch: torch.Tensor,
163
+ poison_ratio: Optional[float] = None,
164
+ ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, List[str]]:
165
+ batch_size = x_batch.size(0)
166
+ new_x = x_batch.clone()
167
+ new_y = y_batch.clone()
168
+ is_poisoned = torch.zeros(batch_size, dtype=torch.bool)
169
+ attack_types = ["clean"] * batch_size
170
+ ratio = poison_ratio if poison_ratio is not None else self.poison_ratio
171
+
172
+ for i in range(batch_size):
173
+ if np.random.rand() < ratio:
174
+ attack = np.random.choice(self.attacks)
175
+ x_p, y_p = attack(x_batch[i], int(y_batch[i].item()))
176
+ new_x[i] = x_p
177
+ new_y[i] = y_p
178
+ is_poisoned[i] = True
179
+ attack_types[i] = attack.name
180
+
181
+ return new_x, new_y, is_poisoned, attack_types
182
+
183
+ def poison_all(
184
+ self, x_batch: torch.Tensor, y_batch: torch.Tensor, attack: BaseAttack
185
+ ) -> Tuple[torch.Tensor, torch.Tensor]:
186
+ """Отруює весь батч однією конкретною атакою — для тестування."""
187
+ new_x = x_batch.clone()
188
+ new_y = y_batch.clone()
189
+ for i in range(x_batch.size(0)):
190
+ x_p, y_p = attack(x_batch[i], int(y_batch[i].item()))
191
+ new_x[i] = x_p
192
+ new_y[i] = y_p
193
+ return new_x, new_y
194
+
195
+
196
+ if __name__ == "__main__":
197
+ # Швидкий тест
198
+ x = torch.rand(8, 3, 32, 32)
199
+ y = torch.randint(0, 10, (8,))
200
+
201
+ gen = PoisonGenerator(poison_ratio=0.5)
202
+ new_x, new_y, mask, types = gen.poison_batch(x, y)
203
+
204
+ print(f"Original shape: {x.shape}")
205
+ print(f"Poisoned mask: {mask.tolist()}")
206
+ print(f"Attack types: {types}")
207
+ print(f"Label changes: {(new_y != y).tolist()}")
208
+ print(f"Pixel diff (max): {(new_x - x).abs().max().item():.4f}")
requirements.txt ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ torch>=2.0.0
2
+ torchvision>=0.15.0
3
+ numpy>=1.24.0
train.py ADDED
@@ -0,0 +1,358 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ Повний пайплайн навчання з захистом від отруєння.
3
+
4
+ Процес:
5
+ Фаза 1: Тренуємо Detector на згенерованих clean/poisoned парах.
6
+ Фаза 2: Тренуємо дві моделі:
7
+ a) Baseline — без захисту, на отруєному датасеті
8
+ b) Protected — з Detector-зважуванням loss
9
+ Фаза 3: Оцінюємо обидві на:
10
+ - Clean test set (звичайна точність)
11
+ - Backdoored test set (наскільки атака спрацьовує — менше = краще)
12
+
13
+ Запуск:
14
+ python train.py --dataset cifar10 --epochs 20 --poison_ratio 0.3
15
+ """
16
+
17
+ import argparse
18
+ import os
19
+ import time
20
+ from typing import Dict, Tuple
21
+
22
+ import torch
23
+ import torch.nn as nn
24
+ import torch.optim as optim
25
+ from torch.utils.data import DataLoader
26
+
27
+ from poison_generator import (
28
+ PoisonGenerator,
29
+ LabelFlipAttack,
30
+ BackdoorAttack,
31
+ CleanLabelAttack,
32
+ FeatureCorruptionAttack,
33
+ )
34
+ from detector import Detector, detector_loss
35
+ from models import ProtectedModel, weighted_cross_entropy
36
+ from hf_data import get_dataloaders, DATASET_CONFIGS
37
+
38
+
39
+ # ---- Мапінг назв атак на ID (для attack_head) ----
40
+ ATTACK_NAME_TO_ID = {
41
+ "clean": 0,
42
+ "label_flip": 1,
43
+ "backdoor": 2,
44
+ "clean_label": 3,
45
+ "feature_corruption": 4,
46
+ }
47
+ NUM_ATTACK_TYPES = len(ATTACK_NAME_TO_ID)
48
+
49
+
50
+ def make_attack_id_tensor(attack_types: list) -> torch.Tensor:
51
+ return torch.tensor([ATTACK_NAME_TO_ID[t] for t in attack_types], dtype=torch.long)
52
+
53
+
54
+ # =============================================================================
55
+ # ФАЗА 1: ТРЕНУВАННЯ DETECTOR'А
56
+ # =============================================================================
57
+ def train_detector(
58
+ detector: Detector,
59
+ train_loader: DataLoader,
60
+ poison_gen: PoisonGenerator,
61
+ epochs: int,
62
+ device: torch.device,
63
+ lr: float = 1e-3,
64
+ ):
65
+ print("\n" + "=" * 60)
66
+ print("ФАЗА 1: Тренування Detector'а")
67
+ print("=" * 60)
68
+
69
+ optimizer = optim.Adam(detector.parameters(), lr=lr)
70
+ detector.train()
71
+
72
+ for epoch in range(epochs):
73
+ epoch_loss = 0.0
74
+ correct = 0
75
+ total = 0
76
+ attack_correct = 0
77
+ start = time.time()
78
+
79
+ for batch_idx, (x, y) in enumerate(train_loader):
80
+ x, y = x.to(device), y.to(device)
81
+
82
+ # Генеруємо отруєний батч (60% отрути для збалансованості при навчанні детектора)
83
+ x_mixed, _, is_poisoned, attack_types = poison_gen.poison_batch(x, y, poison_ratio=0.6)
84
+ x_mixed = x_mixed.to(device)
85
+ is_poisoned = is_poisoned.to(device)
86
+ attack_ids = make_attack_id_tensor(attack_types).to(device)
87
+
88
+ optimizer.zero_grad()
89
+ embed, poison_logits, attack_logits = detector(x_mixed)
90
+
91
+ loss, metrics = detector_loss(
92
+ embed, poison_logits, attack_logits,
93
+ is_poisoned, attack_ids,
94
+ use_contrastive=True,
95
+ )
96
+ loss.backward()
97
+ optimizer.step()
98
+
99
+ epoch_loss += metrics["total"]
100
+ preds = poison_logits.argmax(dim=-1)
101
+ correct += (preds == is_poisoned.long()).sum().item()
102
+ attack_correct += (attack_logits.argmax(dim=-1) == attack_ids).sum().item()
103
+ total += x.size(0)
104
+
105
+ # Оновлюємо memory bank на отруєних зразках
106
+ if is_poisoned.any():
107
+ detector.update_memory(
108
+ embed[is_poisoned].detach(),
109
+ attack_ids[is_poisoned],
110
+ )
111
+
112
+ elapsed = time.time() - start
113
+ print(
114
+ f"Detector epoch {epoch + 1}/{epochs} | "
115
+ f"loss={epoch_loss / len(train_loader):.4f} | "
116
+ f"binary_acc={100 * correct / total:.2f}% | "
117
+ f"attack_type_acc={100 * attack_correct / total:.2f}% | "
118
+ f"time={elapsed:.1f}s | "
119
+ f"memory_size={len(detector.memory_embeds)}"
120
+ )
121
+
122
+
123
+ # =============================================================================
124
+ # ФАЗА 2: ТРЕНУВАННЯ PROTECTED ТА BASELINE МОДЕЛЕЙ
125
+ # =============================================================================
126
+ def train_classifier(
127
+ model: nn.Module,
128
+ detector: Detector, # None для baseline
129
+ train_loader: DataLoader,
130
+ poison_gen: PoisonGenerator,
131
+ epochs: int,
132
+ device: torch.device,
133
+ use_defense: bool,
134
+ lr: float = 1e-3,
135
+ name: str = "Model",
136
+ ):
137
+ print("\n" + "=" * 60)
138
+ print(f"ФАЗА 2: Тренування {name} (захист={'ON' if use_defense else 'OFF'})")
139
+ print("=" * 60)
140
+
141
+ optimizer = optim.Adam(model.parameters(), lr=lr)
142
+ if detector is not None:
143
+ detector.eval()
144
+
145
+ for epoch in range(epochs):
146
+ epoch_loss = 0.0
147
+ correct = 0
148
+ total = 0
149
+ avg_trust_poisoned = 0.0
150
+ avg_trust_clean = 0.0
151
+ num_poisoned = 0
152
+ num_clean = 0
153
+ start = time.time()
154
+
155
+ model.train()
156
+ for batch_idx, (x, y) in enumerate(train_loader):
157
+ x, y = x.to(device), y.to(device)
158
+
159
+ # Отруюємо частину датасету (це симуляція компрометованих даних)
160
+ x_p, y_p, is_poisoned, _ = poison_gen.poison_batch(x, y)
161
+ x_p = x_p.to(device)
162
+ y_p = y_p.to(device)
163
+ is_poisoned = is_poisoned.to(device)
164
+
165
+ optimizer.zero_grad()
166
+ logits = model(x_p)
167
+
168
+ if use_defense and detector is not None:
169
+ # Імунний захист: ваги довіри від детектора
170
+ trust = detector.trust_weights(x_p, soft=True)
171
+ loss = weighted_cross_entropy(logits, y_p, trust)
172
+
173
+ if is_poisoned.any():
174
+ avg_trust_poisoned += trust[is_poisoned].sum().item()
175
+ num_poisoned += is_poisoned.sum().item()
176
+ if (~is_poisoned).any():
177
+ avg_trust_clean += trust[~is_poisoned].sum().item()
178
+ num_clean += (~is_poisoned).sum().item()
179
+ else:
180
+ # Baseline — без захисту
181
+ loss = nn.functional.cross_entropy(logits, y_p)
182
+
183
+ loss.backward()
184
+ optimizer.step()
185
+
186
+ epoch_loss += loss.item()
187
+ correct += (logits.argmax(dim=-1) == y_p).sum().item()
188
+ total += x.size(0)
189
+
190
+ elapsed = time.time() - start
191
+ info = (
192
+ f"{name} epoch {epoch + 1}/{epochs} | "
193
+ f"loss={epoch_loss / len(train_loader):.4f} | "
194
+ f"train_acc(poisoned)={100 * correct / total:.2f}% | "
195
+ f"time={elapsed:.1f}s"
196
+ )
197
+ if use_defense and num_poisoned > 0:
198
+ info += (
199
+ f" | avg_trust(poisoned)={avg_trust_poisoned / num_poisoned:.3f}"
200
+ f" | avg_trust(clean)={avg_trust_clean / max(num_clean, 1):.3f}"
201
+ )
202
+ print(info)
203
+
204
+
205
+ # =============================================================================
206
+ # ФАЗА 3: ОЦІНКА
207
+ # =============================================================================
208
+ @torch.no_grad()
209
+ def evaluate_clean(model: nn.Module, test_loader: DataLoader, device: torch.device) -> float:
210
+ model.eval()
211
+ correct = 0
212
+ total = 0
213
+ for x, y in test_loader:
214
+ x, y = x.to(device), y.to(device)
215
+ logits = model(x)
216
+ correct += (logits.argmax(dim=-1) == y).sum().item()
217
+ total += x.size(0)
218
+ return 100.0 * correct / total
219
+
220
+
221
+ @torch.no_grad()
222
+ def evaluate_backdoor(
223
+ model: nn.Module,
224
+ test_loader: DataLoader,
225
+ backdoor_attack: BackdoorAttack,
226
+ device: torch.device,
227
+ ) -> float:
228
+ """
229
+ Attack Success Rate (ASR): який % НЕ target-class зразків модель класифікує як target після додавання тригера.
230
+ МЕНШЕ = КРАЩЕ. Якщо захист працює — модель не повинна реагувати на тригер.
231
+ """
232
+ model.eval()
233
+ target = backdoor_attack.target_class
234
+ success = 0
235
+ total = 0
236
+
237
+ for x, y in test_loader:
238
+ # Пропускаємо зразки, які вже належать target класу
239
+ mask = y != target
240
+ if mask.sum() == 0:
241
+ continue
242
+
243
+ x_filtered = x[mask]
244
+ # Додаємо тригер
245
+ x_triggered = x_filtered.clone()
246
+ for i in range(x_triggered.size(0)):
247
+ x_triggered[i], _ = backdoor_attack(x_triggered[i], int(y[mask][i].item()))
248
+
249
+ x_triggered = x_triggered.to(device)
250
+ preds = model(x_triggered).argmax(dim=-1)
251
+ success += (preds == target).sum().item()
252
+ total += x_triggered.size(0)
253
+
254
+ return 100.0 * success / max(total, 1)
255
+
256
+
257
+ # =============================================================================
258
+ # MAIN
259
+ # =============================================================================
260
+ def main():
261
+ parser = argparse.ArgumentParser()
262
+ parser.add_argument(
263
+ "--dataset", type=str, default="cifar10",
264
+ choices=list(DATASET_CONFIGS.keys()),
265
+ help="HF dataset: cifar10, cifar100, mnist, tiny_imagenet, fashion_mnist",
266
+ )
267
+ parser.add_argument("--epochs_detector", type=int, default=5)
268
+ parser.add_argument("--epochs_classifier", type=int, default=10)
269
+ parser.add_argument("--batch_size", type=int, default=128)
270
+ parser.add_argument("--poison_ratio", type=float, default=0.3)
271
+ parser.add_argument("--lr", type=float, default=1e-3)
272
+ parser.add_argument("--cache_dir", type=str, default=None, help="HF cache dir (default: ~/.cache/huggingface)")
273
+ parser.add_argument("--save_dir", type=str, default="./checkpoints")
274
+ parser.add_argument("--num_workers", type=int, default=2)
275
+ parser.add_argument("--resize_to", type=int, default=None, help="Resize images to NxN (optional)")
276
+ parser.add_argument("--seed", type=int, default=42)
277
+ args = parser.parse_args()
278
+
279
+ torch.manual_seed(args.seed)
280
+ os.makedirs(args.save_dir, exist_ok=True)
281
+
282
+ device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
283
+ print(f"Device: {device}")
284
+ print(f"Args: {vars(args)}")
285
+
286
+ # Data — тепер через HuggingFace datasets
287
+ train_loader, test_loader, in_channels, num_classes = get_dataloaders(
288
+ args.dataset,
289
+ batch_size=args.batch_size,
290
+ cache_dir=args.cache_dir,
291
+ num_workers=args.num_workers,
292
+ resize_to=args.resize_to,
293
+ )
294
+
295
+ # Poison generator з усіма типами атак
296
+ target_class = 0 # бекдор завжди веде до класу 0
297
+ backdoor = BackdoorAttack(trigger_size=4, trigger_value=1.0, target_class=target_class)
298
+ poison_gen = PoisonGenerator(
299
+ attacks=[
300
+ LabelFlipAttack(num_classes=num_classes),
301
+ backdoor,
302
+ CleanLabelAttack(epsilon=0.05),
303
+ FeatureCorruptionAttack(corruption_ratio=0.2, num_classes=num_classes),
304
+ ],
305
+ poison_ratio=args.poison_ratio,
306
+ num_classes=num_classes,
307
+ )
308
+
309
+ # --- ФАЗА 1: Detector ---
310
+ detector = Detector(
311
+ in_channels=in_channels, embed_dim=128, num_attack_types=NUM_ATTACK_TYPES
312
+ ).to(device)
313
+ train_detector(detector, train_loader, poison_gen, args.epochs_detector, device, args.lr)
314
+ torch.save(detector.state_dict(), os.path.join(args.save_dir, "detector.pt"))
315
+
316
+ # --- ФАЗА 2: Baseline (без захисту) ---
317
+ baseline = ProtectedModel(num_classes=num_classes, in_channels=in_channels).to(device)
318
+ train_classifier(
319
+ baseline, None, train_loader, poison_gen,
320
+ args.epochs_classifier, device, use_defense=False, lr=args.lr, name="BASELINE"
321
+ )
322
+ torch.save(baseline.state_dict(), os.path.join(args.save_dir, "baseline.pt"))
323
+
324
+ # --- ФАЗА 2: Protected (з імунним захистом) ---
325
+ protected = ProtectedModel(num_classes=num_classes, in_channels=in_channels).to(device)
326
+ train_classifier(
327
+ protected, detector, train_loader, poison_gen,
328
+ args.epochs_classifier, device, use_defense=True, lr=args.lr, name="PROTECTED"
329
+ )
330
+ torch.save(protected.state_dict(), os.path.join(args.save_dir, "protected.pt"))
331
+
332
+ # --- ФАЗА 3: ОЦІНКА ---
333
+ print("\n" + "=" * 60)
334
+ print("ФАЗА 3: Фінальна оцінка")
335
+ print("=" * 60)
336
+
337
+ baseline_clean = evaluate_clean(baseline, test_loader, device)
338
+ protected_clean = evaluate_clean(protected, test_loader, device)
339
+ baseline_asr = evaluate_backdoor(baseline, test_loader, backdoor, device)
340
+ protected_asr = evaluate_backdoor(protected, test_loader, backdoor, device)
341
+
342
+ print(f"\n{'Метрика':<35} {'Baseline':<15} {'Protected':<15}")
343
+ print("-" * 65)
344
+ print(f"{'Clean accuracy ↑':<35} {baseline_clean:<15.2f} {protected_clean:<15.2f}")
345
+ print(f"{'Backdoor ASR ↓ (атака успішна %)':<35} {baseline_asr:<15.2f} {protected_asr:<15.2f}")
346
+
347
+ print("\nІнтерпретація:")
348
+ print(f" • Clean accuracy: вища = краще (нормальна продуктивність)")
349
+ print(f" • Backdoor ASR: нижча = краще (атака менш ефективна)")
350
+ if protected_asr < baseline_asr:
351
+ delta = baseline_asr - protected_asr
352
+ print(f" ✓ Захист знизив успішність атаки на {delta:.2f} процентних пунктів!")
353
+ else:
354
+ print(f" ✗ Захист не зменшив атаку — треба тюнити Detector.")
355
+
356
+
357
+ if __name__ == "__main__":
358
+ main()