Upload 16 files
Browse files- .gitattributes +6 -0
- checkpoints/detector.pt +3 -0
- data/cifar-10-batches-py/batches.meta +0 -0
- data/cifar-10-batches-py/data_batch_1 +3 -0
- data/cifar-10-batches-py/data_batch_2 +3 -0
- data/cifar-10-batches-py/data_batch_3 +3 -0
- data/cifar-10-batches-py/data_batch_4 +3 -0
- data/cifar-10-batches-py/data_batch_5 +3 -0
- data/cifar-10-batches-py/readme.html +1 -0
- data/cifar-10-batches-py/test_batch +3 -0
- data/cifar-10-python.tar.gz +3 -0
- detector.py +230 -0
- hf_data.py +185 -0
- models.py +102 -0
- poison_generator.py +208 -0
- requirements.txt +3 -0
- train.py +358 -0
.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()
|