kcai-data-sampling 0.0.3__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- kcai_data_sampling/__init__.py +9 -0
- kcai_data_sampling/adversarial_run.py +524 -0
- kcai_data_sampling/augment_run.py +807 -0
- kcai_data_sampling/kc_inference.py +357 -0
- kcai_data_sampling/run_augment_pipeline.py +658 -0
- kcai_data_sampling-0.0.3.dist-info/METADATA +83 -0
- kcai_data_sampling-0.0.3.dist-info/RECORD +9 -0
- kcai_data_sampling-0.0.3.dist-info/WHEEL +5 -0
- kcai_data_sampling-0.0.3.dist-info/top_level.txt +1 -0
|
@@ -0,0 +1,524 @@
|
|
|
1
|
+
# adversarial_runner.py
|
|
2
|
+
from __future__ import annotations
|
|
3
|
+
import os, io, json, uuid, hashlib, random, pathlib, base64, time, math, shutil
|
|
4
|
+
from typing import Any, Dict, List, Optional, Tuple
|
|
5
|
+
|
|
6
|
+
import numpy as np
|
|
7
|
+
import pandas as pd
|
|
8
|
+
from PIL import Image
|
|
9
|
+
import yaml
|
|
10
|
+
import cv2
|
|
11
|
+
import pyarrow as pa
|
|
12
|
+
import pyarrow.parquet as pq
|
|
13
|
+
import torch
|
|
14
|
+
import torch.nn.functional as F
|
|
15
|
+
from tqdm import tqdm
|
|
16
|
+
|
|
17
|
+
# attaques sous-jacentes (utilisés dans robust-ai)
|
|
18
|
+
# CleverHans pour FGSM/PGD/CW
|
|
19
|
+
|
|
20
|
+
from cleverhans.torch.attacks.fast_gradient_method import fast_gradient_method # FGSM
|
|
21
|
+
from cleverhans.torch.attacks.projected_gradient_descent import projected_gradient_descent # PGD
|
|
22
|
+
from cleverhans.torch.attacks.carlini_wagner_l2 import carlini_wagner_l2 # C&W (L2)
|
|
23
|
+
# AutoAttack / APGD
|
|
24
|
+
from autoattack import AutoAttack
|
|
25
|
+
|
|
26
|
+
# def _seed32(*parts, base_seed: int = 0) -> int:
|
|
27
|
+
# h = hashlib.blake2b(digest_size=8)
|
|
28
|
+
# for p in parts:
|
|
29
|
+
# h.update(str(p).encode("utf-8")); h.update(b"|")
|
|
30
|
+
# return (int.from_bytes(h.digest(), "little") ^ (base_seed & 0xFFFFFFFF)) & 0xFFFFFFFF
|
|
31
|
+
|
|
32
|
+
def _ensure_dir(path: str):
|
|
33
|
+
if path:
|
|
34
|
+
pathlib.Path(path).mkdir(parents=True, exist_ok=True)
|
|
35
|
+
|
|
36
|
+
def _pil_from_bytes(b: bytes) -> Image.Image:
|
|
37
|
+
return Image.open(io.BytesIO(b)).convert("RGB")
|
|
38
|
+
|
|
39
|
+
def _np_from_pil(img: Image.Image) -> np.ndarray:
|
|
40
|
+
return np.array(img)
|
|
41
|
+
|
|
42
|
+
def _encode_image(arr_rgb: np.ndarray, ext: str = "png") -> bytes:
|
|
43
|
+
if ext.lower() in ("jpg","jpeg"):
|
|
44
|
+
ok, buf = cv2.imencode(".jpg", cv2.cvtColor(arr_rgb, cv2.COLOR_RGB2BGR), [int(cv2.IMWRITE_JPEG_QUALITY), 95])
|
|
45
|
+
elif ext.lower() in ("tif","tiff"):
|
|
46
|
+
ok, buf = cv2.imencode(".tif", cv2.cvtColor(arr_rgb, cv2.COLOR_RGB2BGR))
|
|
47
|
+
else:
|
|
48
|
+
ok, buf = cv2.imencode(".png", cv2.cvtColor(arr_rgb, cv2.COLOR_RGB2BGR))
|
|
49
|
+
if not ok: raise RuntimeError("Échec encodage image")
|
|
50
|
+
return buf.tobytes()
|
|
51
|
+
|
|
52
|
+
def _load_image_from_row(row: pd.Series, image_source: str, col_image_bytes: Optional[str], col_file_path: Optional[str], base_dir: Optional[str]=None) -> np.ndarray:
|
|
53
|
+
if image_source == "image_bytes":
|
|
54
|
+
b = row[col_image_bytes]
|
|
55
|
+
if isinstance(b, str):
|
|
56
|
+
try: b = b.encode("latin1")
|
|
57
|
+
except Exception: b = base64.b64decode(b)
|
|
58
|
+
return _np_from_pil(_pil_from_bytes(b))
|
|
59
|
+
fp = str(row[col_file_path])
|
|
60
|
+
if base_dir and not os.path.isabs(fp): fp = os.path.join(base_dir, fp)
|
|
61
|
+
arr = cv2.imread(fp, cv2.IMREAD_UNCHANGED)
|
|
62
|
+
if arr is None: raise FileNotFoundError(fp)
|
|
63
|
+
if arr.ndim == 2: arr = cv2.cvtColor(arr, cv2.COLOR_GRAY2RGB)
|
|
64
|
+
else: arr = cv2.cvtColor(arr, cv2.COLOR_BGR2RGB)
|
|
65
|
+
return arr
|
|
66
|
+
|
|
67
|
+
def _write_table_single(df: pd.DataFrame, table_out_path: str, out_type: str, compression: Optional[str]):
|
|
68
|
+
_ensure_dir(os.path.dirname(table_out_path))
|
|
69
|
+
if out_type == "csv": df.to_csv(table_out_path, index=False)
|
|
70
|
+
elif out_type == "parquet": df.to_parquet(table_out_path, index=False, compression=compression)
|
|
71
|
+
else: raise ValueError(f"Unsupported table output type: {out_type}")
|
|
72
|
+
|
|
73
|
+
def _clear_partition_dir(root_path: str, partition_cols: List[str], key_values: List[str]):
|
|
74
|
+
path = root_path
|
|
75
|
+
for col, val in zip(partition_cols, key_values):
|
|
76
|
+
path = os.path.join(path, f"{col}={val}")
|
|
77
|
+
shutil.rmtree(path, ignore_errors=True)
|
|
78
|
+
|
|
79
|
+
def _write_parquet_dataset(
|
|
80
|
+
df: pd.DataFrame, root_path: str, partition_cols: List[str], compression: Optional[str],
|
|
81
|
+
existing_data_behavior: str="overwrite_or_ignore", max_rows_per_file: Optional[int]=None,
|
|
82
|
+
files_per_partition_min: Optional[int]=None, max_rows_per_group: Optional[int]=None,
|
|
83
|
+
):
|
|
84
|
+
if df.empty:
|
|
85
|
+
_ensure_dir(root_path); return
|
|
86
|
+
df2 = df.copy()
|
|
87
|
+
if "transform_level" in partition_cols:
|
|
88
|
+
df2["transform_level"] = df2["transform_level"].apply(lambda v: "none" if pd.isna(v) else str(v))
|
|
89
|
+
for col in partition_cols:
|
|
90
|
+
if col != "transform_level":
|
|
91
|
+
df2[col] = df2[col].astype(str)
|
|
92
|
+
_ensure_dir(root_path)
|
|
93
|
+
mrf = int(max_rows_per_file) if max_rows_per_file else None
|
|
94
|
+
mrg = int(max_rows_per_group) if max_rows_per_group else None
|
|
95
|
+
if mrf and (not mrg or mrg > mrf): mrg = mrf
|
|
96
|
+
if not files_per_partition_min or int(files_per_partition_min) <= 1:
|
|
97
|
+
table = pa.Table.from_pandas(df2, preserve_index=False)
|
|
98
|
+
kwargs = {}
|
|
99
|
+
if mrf: kwargs["max_rows_per_file"] = mrf
|
|
100
|
+
if mrg: kwargs["max_rows_per_group"] = mrg; kwargs["row_group_size"] = mrg
|
|
101
|
+
pq.write_to_dataset(table, root_path=root_path, partition_cols=partition_cols,
|
|
102
|
+
compression=compression if compression else None,
|
|
103
|
+
existing_data_behavior=existing_data_behavior, **kwargs)
|
|
104
|
+
return
|
|
105
|
+
grouped = df2.groupby(partition_cols, dropna=False, sort=False)
|
|
106
|
+
for key_vals, part in grouped:
|
|
107
|
+
key_vals = key_vals if isinstance(key_vals, tuple) else (key_vals,)
|
|
108
|
+
key_vals_str = [str(v) for v in key_vals]
|
|
109
|
+
if existing_data_behavior == "delete_matching":
|
|
110
|
+
_clear_partition_dir(root_path, partition_cols, key_vals_str)
|
|
111
|
+
nrows = len(part)
|
|
112
|
+
by_size = math.ceil(nrows / mrf) if mrf else 1
|
|
113
|
+
n_files = max(int(files_per_partition_min), by_size)
|
|
114
|
+
rows_per_shard = max(1, math.ceil(nrows / n_files))
|
|
115
|
+
for i in range(0, nrows, rows_per_shard):
|
|
116
|
+
shard = part.iloc[i:i+rows_per_shard]
|
|
117
|
+
table = pa.Table.from_pandas(shard, preserve_index=False)
|
|
118
|
+
pq.write_to_dataset(table, root_path=root_path, partition_cols=partition_cols,
|
|
119
|
+
compression=compression if compression else None,
|
|
120
|
+
existing_data_behavior="overwrite_or_ignore")
|
|
121
|
+
|
|
122
|
+
def _stable_uid(base_id: str, tname: str, level: Any, params: Dict[str, Any]) -> str:
|
|
123
|
+
h = hashlib.sha1()
|
|
124
|
+
h.update(json.dumps({"base": base_id, "tname": tname, "level": level, "params": params}, sort_keys=True).encode("utf-8"))
|
|
125
|
+
return h.hexdigest()[:10]
|
|
126
|
+
|
|
127
|
+
class SafeDict(dict):
|
|
128
|
+
def __missing__(self, key): return ""
|
|
129
|
+
|
|
130
|
+
# modèle & préproc
|
|
131
|
+
def _to_tensor(img: np.ndarray, resize: Optional[Tuple[int,int]]=None) -> torch.Tensor:
|
|
132
|
+
if resize: img = cv2.resize(img, (resize[1], resize[0]), interpolation=cv2.INTER_LINEAR)
|
|
133
|
+
t = torch.from_numpy(img).float()/255.0 # [0,1]
|
|
134
|
+
t = t.permute(2,0,1) # HWC->CHW
|
|
135
|
+
return t
|
|
136
|
+
|
|
137
|
+
def _normalize(t: torch.Tensor, mean, std):
|
|
138
|
+
mean = torch.tensor(mean, device=t.device).view(3,1,1)
|
|
139
|
+
std = torch.tensor(std, device=t.device).view(3,1,1)
|
|
140
|
+
return (t - mean) / std
|
|
141
|
+
|
|
142
|
+
def _denormalize(t: torch.Tensor, mean, std):
|
|
143
|
+
mean = torch.tensor(mean, device=t.device).view(3,1,1)
|
|
144
|
+
std = torch.tensor(std, device=t.device).view(3,1,1)
|
|
145
|
+
return (t * std + mean)
|
|
146
|
+
|
|
147
|
+
def _clip01(t: torch.Tensor) -> torch.Tensor:
|
|
148
|
+
return torch.clamp(t, 0.0, 1.0)
|
|
149
|
+
|
|
150
|
+
def _predict_logits(model, x: torch.Tensor) -> torch.Tensor:
|
|
151
|
+
model.eval()
|
|
152
|
+
with torch.no_grad():
|
|
153
|
+
return model(x)
|
|
154
|
+
|
|
155
|
+
# adv attacks
|
|
156
|
+
def run_fgsm(model, x, y, eps, clip_min=-1.0, clip_max=1.0):
|
|
157
|
+
# CleverHans attend x normalisé comme utilisé au forward
|
|
158
|
+
def f(x_): return model(x_)
|
|
159
|
+
x_adv = fast_gradient_method(model_fn=f, x=x, eps=eps, norm=np.inf, y=y, targeted=False, clip_min=clip_min, clip_max=clip_max)
|
|
160
|
+
return x_adv
|
|
161
|
+
|
|
162
|
+
def run_pgd(model, x, y, eps, alpha, steps, rand_init=True, clip_min=-1.0, clip_max=1.0):
|
|
163
|
+
def f(x_): return model(x_)
|
|
164
|
+
x_adv = projected_gradient_descent(model_fn=f, x=x, eps=eps, eps_iter=alpha, nb_iter=steps,
|
|
165
|
+
norm=np.inf, y=y, targeted=False, rand_init=rand_init,
|
|
166
|
+
clip_min=clip_min, clip_max=clip_max)
|
|
167
|
+
return x_adv
|
|
168
|
+
|
|
169
|
+
def run_cw(model, x, y, steps=1000, c=0.1, k=0.0, lr=5e-3, clip_min=-1.0, clip_max=1.0):
|
|
170
|
+
# C&W L2 dans CleverHans; renvoie un tenseur adv
|
|
171
|
+
# Get number of classes from model's final layer
|
|
172
|
+
if hasattr(model, 'fc') and hasattr(model.fc, 'out_features'):
|
|
173
|
+
n_classes = model.fc.out_features
|
|
174
|
+
elif hasattr(model, 'classifier') and hasattr(model.classifier, 'out_features'):
|
|
175
|
+
n_classes = model.classifier.out_features
|
|
176
|
+
else:
|
|
177
|
+
n_classes = 1000 # fallback
|
|
178
|
+
|
|
179
|
+
x_adv = carlini_wagner_l2(model, x, n_classes=n_classes,
|
|
180
|
+
targeted=False, y=y, lr=lr, confidence=k,
|
|
181
|
+
clip_min=clip_min, clip_max=clip_max,
|
|
182
|
+
max_iterations=steps, initial_const=c)
|
|
183
|
+
return x_adv
|
|
184
|
+
|
|
185
|
+
def run_apgd(model, x, y, eps, n_restarts=1):
|
|
186
|
+
# AutoAttack s'utilise en lot; on passe un wrapper
|
|
187
|
+
# Utiliser le device des données d'entrée (x.device) au lieu de forcer CUDA
|
|
188
|
+
device = x.device
|
|
189
|
+
aa = AutoAttack(model, norm='Linf', eps=eps, version='custom', verbose=False, device=device)
|
|
190
|
+
aa.attacks_to_run = ['apgd-ce']
|
|
191
|
+
aa.n_restarts = int(n_restarts)
|
|
192
|
+
x_adv = aa.run_standard_evaluation(x, y, bs=x.shape[0], return_labels=False)
|
|
193
|
+
return x_adv
|
|
194
|
+
|
|
195
|
+
# runner principal
|
|
196
|
+
def run_pipeline(cfg: Dict[str, Any]) -> str:
|
|
197
|
+
ap = cfg["adversarial_pipeline"]
|
|
198
|
+
|
|
199
|
+
# model
|
|
200
|
+
mcfg = ap["model"]
|
|
201
|
+
resize_hw = mcfg.get("resize")
|
|
202
|
+
mean = mcfg.get("mean", [0.0,0.0,0.0])
|
|
203
|
+
std = mcfg.get("std", [1.0,1.0,1.0])
|
|
204
|
+
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
|
205
|
+
if device.type == "cuda":
|
|
206
|
+
try:
|
|
207
|
+
_ = torch.tensor([1.0], device=device) + 1.0
|
|
208
|
+
torch.cuda.synchronize()
|
|
209
|
+
except Exception as e:
|
|
210
|
+
print(f"CUDA non utilisable avec cette build ({str(e)[:100]}). Bascule sur CPU.")
|
|
211
|
+
device = torch.device("cpu")
|
|
212
|
+
ckpt = mcfg["ckpt_path"]
|
|
213
|
+
model_name = os.path.basename(ckpt) # nom du modèle
|
|
214
|
+
# Ensure quantized backend is set on CPU BEFORE loading TorchScript
|
|
215
|
+
if torch.device("cpu").type == 'cpu':
|
|
216
|
+
supported_engines = torch.backends.quantized.supported_engines
|
|
217
|
+
if 'fbgemm' in supported_engines:
|
|
218
|
+
torch.backends.quantized.engine = 'fbgemm'
|
|
219
|
+
elif 'qnnpack' in supported_engines:
|
|
220
|
+
torch.backends.quantized.engine = 'qnnpack'
|
|
221
|
+
# Note: we don't print here to keep logs concise
|
|
222
|
+
|
|
223
|
+
# Try to load as TorchScript first (for .pt files), otherwise as regular checkpoint
|
|
224
|
+
try:
|
|
225
|
+
# Force TorchScript quantized models to CPU
|
|
226
|
+
model = torch.jit.load(ckpt, map_location=torch.device("cpu"))
|
|
227
|
+
print(f"✅ Modèle TorchScript chargé : {model_name} (CPU)")
|
|
228
|
+
device = torch.device("cpu")
|
|
229
|
+
except RuntimeError:
|
|
230
|
+
# Not a TorchScript file, try regular checkpoint on CPU
|
|
231
|
+
checkpoint = torch.load(ckpt, map_location=torch.device("cpu"), weights_only=False)
|
|
232
|
+
device = torch.device("cpu")
|
|
233
|
+
if isinstance(checkpoint, dict):
|
|
234
|
+
if "model_state_dict" in checkpoint:
|
|
235
|
+
# Standard checkpoint format with model state dict
|
|
236
|
+
import torchvision.models as models
|
|
237
|
+
arch = checkpoint.get("arch", "resnet18")
|
|
238
|
+
num_classes = checkpoint.get("num_classes", 1000)
|
|
239
|
+
if arch == "resnet18":
|
|
240
|
+
model = models.resnet18(pretrained=False)
|
|
241
|
+
model.fc = torch.nn.Linear(model.fc.in_features, num_classes)
|
|
242
|
+
else:
|
|
243
|
+
raise ValueError(f"Unsupported architecture: {arch}")
|
|
244
|
+
model.load_state_dict(checkpoint["model_state_dict"])
|
|
245
|
+
else:
|
|
246
|
+
model = checkpoint
|
|
247
|
+
else:
|
|
248
|
+
model = checkpoint
|
|
249
|
+
print(f"✅ Modèle checkpoint chargé : {model_name} (CPU)")
|
|
250
|
+
|
|
251
|
+
if hasattr(model, "eval"): model.eval()
|
|
252
|
+
model = model.to(device)
|
|
253
|
+
|
|
254
|
+
# data - import
|
|
255
|
+
dl = ap["dataloader"]
|
|
256
|
+
full_df = pd.read_parquet(dl["path"]) if dl["type"]=="parquet" else pd.read_csv(dl["path"])
|
|
257
|
+
cols = dl["columns"]
|
|
258
|
+
image_source = cols.get("image_source", "image_bytes")
|
|
259
|
+
col_bytes = cols.get("image_bytes"); col_path = cols.get("file_path")
|
|
260
|
+
batch_size = int(dl.get("batch_size", 128))
|
|
261
|
+
|
|
262
|
+
# data - output
|
|
263
|
+
out = ap["outputs"]
|
|
264
|
+
file_ext = out["image_output"].get("file_ext", "png")
|
|
265
|
+
t_out = out["table_output"]; m_out = out.get("metrics_output")
|
|
266
|
+
ordered_front = out["columns"]["ordered_front"]
|
|
267
|
+
include_src = out["columns"].get("include_source_columns", False)
|
|
268
|
+
|
|
269
|
+
# run opts
|
|
270
|
+
run_opts = ap.get("run", {})
|
|
271
|
+
seed = int(run_opts.get("seed", 123))
|
|
272
|
+
deterministic_ids = bool(run_opts.get("deterministic_ids", False))
|
|
273
|
+
id_pattern = run_opts.get("id_pattern", "{base}__{tname}__{level}__{uid}")
|
|
274
|
+
output_sample_type = run_opts.get("output_sample_type", "adversarial")
|
|
275
|
+
rng = np.random.RandomState(seed)
|
|
276
|
+
|
|
277
|
+
attacks = ap["attacks"]
|
|
278
|
+
all_records: List[Dict[str, Any]] = [] # Pour les métriques finales
|
|
279
|
+
|
|
280
|
+
# Calcul des bornes de clipping dans l'espace normalisé
|
|
281
|
+
# Si input_range est [0,1] et on applique (x - mean) / std :
|
|
282
|
+
# clip_min = (0 - mean) / std, clip_max = (1 - mean) / std
|
|
283
|
+
clip_min_normalized = (0.0 - np.array(mean)) / np.array(std)
|
|
284
|
+
clip_max_normalized = (1.0 - np.array(mean)) / np.array(std)
|
|
285
|
+
clip_min = float(clip_min_normalized.min()) # Prendre le min sur les 3 canaux
|
|
286
|
+
clip_max = float(clip_max_normalized.max()) # Prendre le max sur les 3 canaux
|
|
287
|
+
|
|
288
|
+
def _sample_level(a_cfg):
|
|
289
|
+
sp = a_cfg.get("sampling")
|
|
290
|
+
if not sp: return None
|
|
291
|
+
lo, hi = map(float, sp["range"])
|
|
292
|
+
val = float(rng.uniform(lo, hi))
|
|
293
|
+
rd = sp.get("round");
|
|
294
|
+
if rd is not None: val = round(val, int(rd))
|
|
295
|
+
return val
|
|
296
|
+
|
|
297
|
+
n = len(full_df)
|
|
298
|
+
num_batches = (n + batch_size - 1) // batch_size
|
|
299
|
+
total_operations = num_batches * len(attacks)
|
|
300
|
+
|
|
301
|
+
print(f"\n{'='*60}")
|
|
302
|
+
print(f" --- Starting adversarial pipeline ---")
|
|
303
|
+
print(f" 📊 {n} images | {len(attacks)} attacks | {num_batches} batches")
|
|
304
|
+
print(f" ⚙️ Total: {total_operations} operations")
|
|
305
|
+
print(f"{'='*60}\n")
|
|
306
|
+
|
|
307
|
+
pbar = tqdm(total=total_operations, desc="Global pipeline", unit="op", ncols=100)
|
|
308
|
+
|
|
309
|
+
# loop on attacks (for incremental writing)
|
|
310
|
+
for attack_idx, a_cfg in enumerate(attacks):
|
|
311
|
+
alias = a_cfg.get("alias", a_cfg["name"])
|
|
312
|
+
aid = int(a_cfg.get("id", 0))
|
|
313
|
+
|
|
314
|
+
print(f"\n🎯 Attack {attack_idx+1}/{len(attacks)}: {alias} (id={aid})")
|
|
315
|
+
|
|
316
|
+
records: List[Dict[str, Any]] = [] # Records for this attack
|
|
317
|
+
|
|
318
|
+
# loop on batch for this attack
|
|
319
|
+
for start in range(0, n, batch_size):
|
|
320
|
+
part = full_df.iloc[start:start+batch_size]
|
|
321
|
+
# prepare batch tensors
|
|
322
|
+
imgs, meta = [], []
|
|
323
|
+
for _, row in part.iterrows():
|
|
324
|
+
try:
|
|
325
|
+
arr = _load_image_from_row(row, image_source, col_bytes, col_path)
|
|
326
|
+
except Exception:
|
|
327
|
+
continue
|
|
328
|
+
t = _to_tensor(arr, resize=resize_hw)
|
|
329
|
+
imgs.append(t); meta.append(row)
|
|
330
|
+
if not imgs:
|
|
331
|
+
pbar.update(1)
|
|
332
|
+
continue
|
|
333
|
+
x0 = torch.stack(imgs, dim=0).to(device) # [B,3,H,W] in [0,1]
|
|
334
|
+
x_n = _normalize(x0, mean, std)
|
|
335
|
+
|
|
336
|
+
# reference labels (if class_id available); otherwise logits argmax of the model
|
|
337
|
+
with torch.no_grad():
|
|
338
|
+
logits = model(x_n)
|
|
339
|
+
if "class_id" in cols and cols["class_id"] in part.columns and part[cols["class_id"]].notna().all():
|
|
340
|
+
y = torch.tensor(part[cols["class_id"]].values, device=device).long()
|
|
341
|
+
else:
|
|
342
|
+
y = logits.argmax(1)
|
|
343
|
+
|
|
344
|
+
level = _sample_level(a_cfg)
|
|
345
|
+
params = dict(a_cfg.get("params", {}))
|
|
346
|
+
lvl_param = a_cfg.get("level_param");
|
|
347
|
+
if lvl_param is not None and level is not None: params[lvl_param] = level
|
|
348
|
+
|
|
349
|
+
# update progress bar description
|
|
350
|
+
batch_idx = start // batch_size + 1
|
|
351
|
+
level_str = f"lvl={level:.3f}" if level is not None else "lvl=N/A"
|
|
352
|
+
pbar.set_description(f"Batch {batch_idx}/{num_batches} | 🎯 {alias} ({level_str})")
|
|
353
|
+
|
|
354
|
+
# prepare x for attack (normalized)
|
|
355
|
+
x_in = x_n.detach()
|
|
356
|
+
|
|
357
|
+
# dispatch attacks
|
|
358
|
+
lib = a_cfg.get("lib", "cleverhans").lower()
|
|
359
|
+
if alias == "fgsm" and lib == "cleverhans":
|
|
360
|
+
eps = float(params["eps"])
|
|
361
|
+
x_adv_n = run_fgsm(model, x_in, y, eps=eps, clip_min=clip_min, clip_max=clip_max)
|
|
362
|
+
elif alias == "pgd" and lib == "cleverhans":
|
|
363
|
+
eps = float(params["eps"])
|
|
364
|
+
alpha = float(a_cfg.get("step_size_ratio", 0.25)) * eps
|
|
365
|
+
steps = int(params.get("steps", 10))
|
|
366
|
+
x_adv_n = run_pgd(model, x_in, y, eps=eps, alpha=alpha, steps=steps, rand_init=bool(params.get("rand_init", True)), clip_min=clip_min, clip_max=clip_max)
|
|
367
|
+
elif alias == "cw" and lib == "cleverhans":
|
|
368
|
+
steps = int(params.get("steps", 1000))
|
|
369
|
+
c = float(params.get("c", 0.1))
|
|
370
|
+
k = float(params.get("confidence", 0.0))
|
|
371
|
+
lr_cw = float(params.get("lr", 5e-3))
|
|
372
|
+
# C&W expects normalized input and works in the normalized space
|
|
373
|
+
x_adv_n = run_cw(model, x_in, y, steps=steps, c=c, k=k, lr=lr_cw, clip_min=clip_min, clip_max=clip_max)
|
|
374
|
+
elif alias.startswith("apgd") and lib == "autoattack":
|
|
375
|
+
eps = float(params["eps"])
|
|
376
|
+
n_restarts = int(params.get("n_restarts", 1))
|
|
377
|
+
# AutoAttack returns already preprocessed pixels for the model → here it's in the normalized space
|
|
378
|
+
x_adv_n = run_apgd(model, x_in, y, eps=eps, n_restarts=n_restarts)
|
|
379
|
+
else:
|
|
380
|
+
raise ValueError(f"Attack not supported: {alias} ({lib})")
|
|
381
|
+
|
|
382
|
+
# denormalize to save PNG
|
|
383
|
+
x_adv = _clip01(_denormalize(x_adv_n, mean, std))
|
|
384
|
+
# NOTE: Inference removed - use kc_inference.py to add predictions
|
|
385
|
+
|
|
386
|
+
# save lines
|
|
387
|
+
for i in range(x_adv.shape[0]):
|
|
388
|
+
arr = (x_adv[i].detach().cpu().permute(1,2,0).numpy()*255.0).round().astype(np.uint8)
|
|
389
|
+
img_bytes = _encode_image(arr, file_ext)
|
|
390
|
+
|
|
391
|
+
# Calculate hash of the adversarial image
|
|
392
|
+
img_hash = hashlib.sha256(img_bytes).hexdigest()
|
|
393
|
+
|
|
394
|
+
# Generate UUIDs
|
|
395
|
+
sample_uuid_val = str(uuid.uuid4())
|
|
396
|
+
parent_sample_uuid_val = meta[i].get(cols.get("sample_uuid", "sample_uuid"), str(uuid.uuid4()))
|
|
397
|
+
|
|
398
|
+
base_id = meta[i].get(cols.get("sample_id","sample_id"), f"row{start+i}")
|
|
399
|
+
params_used = {"alias": alias, **params}
|
|
400
|
+
uid = _stable_uid(base_id, alias, level, params_used) if deterministic_ids else str(uuid.uuid4())[:8]
|
|
401
|
+
child_id = (run_opts.get("id_pattern","{base}__{tname}__{level}__{uid}")).format_map(
|
|
402
|
+
SafeDict({"base": base_id, "tname": alias, "level": ("none" if level is None else str(level)), "uid": uid})
|
|
403
|
+
)
|
|
404
|
+
rec = {
|
|
405
|
+
"sample_uuid": sample_uuid_val,
|
|
406
|
+
"parent_sample_uuid": parent_sample_uuid_val,
|
|
407
|
+
"sample_id": base_id,
|
|
408
|
+
"sample_type": output_sample_type,
|
|
409
|
+
"class_id": int(meta[i].get(cols.get("class_id","class_id"), -1)) if cols.get("class_id") in meta[i] else None,
|
|
410
|
+
"class_name": meta[i].get(cols.get("class_name","class_name"), None),
|
|
411
|
+
"split": meta[i].get(cols.get("split","split"), None),
|
|
412
|
+
"transform_id": aid,
|
|
413
|
+
"transform_name": alias,
|
|
414
|
+
"transform_level": level,
|
|
415
|
+
"transform_params": json.dumps(params_used, ensure_ascii=False),
|
|
416
|
+
"image_bytes": img_bytes,
|
|
417
|
+
"image_hash": img_hash,
|
|
418
|
+
}
|
|
419
|
+
if include_src:
|
|
420
|
+
for c in full_df.columns:
|
|
421
|
+
if c not in rec: rec[f"src__{c}"] = meta[i].get(c, None)
|
|
422
|
+
records.append(rec)
|
|
423
|
+
|
|
424
|
+
# update progress bar after each batch
|
|
425
|
+
pbar.update(1)
|
|
426
|
+
|
|
427
|
+
# incremental writing after each attack
|
|
428
|
+
if records:
|
|
429
|
+
attack_df = pd.DataFrame(records)
|
|
430
|
+
if not attack_df.empty:
|
|
431
|
+
front = ordered_front
|
|
432
|
+
rest = [c for c in attack_df.columns if c not in front]
|
|
433
|
+
attack_df = attack_df[front + rest]
|
|
434
|
+
|
|
435
|
+
# write data for this attack
|
|
436
|
+
if out["table_output"].get("dataset", False) and out["table_output"].get("partition_by"):
|
|
437
|
+
# Write as partitioned dataset
|
|
438
|
+
_write_parquet_dataset(
|
|
439
|
+
attack_df,
|
|
440
|
+
out["table_output"]["path"],
|
|
441
|
+
out["table_output"].get("partition_by", ["transform_name","transform_level"]),
|
|
442
|
+
out["table_output"].get("compression"),
|
|
443
|
+
existing_data_behavior="overwrite_or_ignore",
|
|
444
|
+
max_rows_per_file=out["table_output"].get("max_rows_per_file"),
|
|
445
|
+
files_per_partition_min=out["table_output"].get("files_per_partition_min"),
|
|
446
|
+
max_rows_per_group=out["table_output"].get("max_rows_per_group"),
|
|
447
|
+
)
|
|
448
|
+
print(f" ✅ {len(records)} échantillons écrits pour {alias} (partitioned)")
|
|
449
|
+
else:
|
|
450
|
+
# For single file mode, accumulate all records and write at the end
|
|
451
|
+
print(f" 📦 {len(records)} échantillons accumulés pour {alias}")
|
|
452
|
+
|
|
453
|
+
all_records.extend(records)
|
|
454
|
+
|
|
455
|
+
# free memory
|
|
456
|
+
del records, attack_df
|
|
457
|
+
|
|
458
|
+
torch.cuda.empty_cache() if torch.cuda.is_available() else None
|
|
459
|
+
|
|
460
|
+
pbar.close()
|
|
461
|
+
|
|
462
|
+
print(f"\n{'='*60}")
|
|
463
|
+
print(f"✅ Pipeline terminated | {len(all_records)} samples generated")
|
|
464
|
+
print(f"{'='*60}\n")
|
|
465
|
+
|
|
466
|
+
# The data has already been written incrementally (if dataset mode)
|
|
467
|
+
# For single file mode, write all accumulated records now
|
|
468
|
+
out_df = pd.DataFrame(all_records)
|
|
469
|
+
|
|
470
|
+
if not out["table_output"].get("dataset", False) or not out["table_output"].get("partition_by"):
|
|
471
|
+
# Single file mode: write all records to one file
|
|
472
|
+
if not out_df.empty:
|
|
473
|
+
front = out["columns"]["ordered_front"]
|
|
474
|
+
rest = [c for c in out_df.columns if c not in front]
|
|
475
|
+
out_df = out_df[front + rest]
|
|
476
|
+
|
|
477
|
+
output_path = out["table_output"]["path"]
|
|
478
|
+
compression = out["table_output"].get("compression", "snappy")
|
|
479
|
+
_write_table_single(out_df, output_path, "parquet", compression)
|
|
480
|
+
print(f"✅ Wrote {len(out_df)} records to single file: {output_path}")
|
|
481
|
+
else:
|
|
482
|
+
print("⚠️ No records to write")
|
|
483
|
+
|
|
484
|
+
# aggregated metrics (by attack/level)
|
|
485
|
+
# NOTE: success_rate removed - use kc_inference.py to calculate after adding predictions
|
|
486
|
+
if out.get("metrics_output"):
|
|
487
|
+
rows = []
|
|
488
|
+
if not out_df.empty:
|
|
489
|
+
g = out_df.groupby(["transform_id","transform_name"], dropna=False)
|
|
490
|
+
for (tid, tname), df_g in g:
|
|
491
|
+
rows.append({
|
|
492
|
+
"scope": "attack",
|
|
493
|
+
"transform_id": int(tid) if pd.notna(tid) else None,
|
|
494
|
+
"transform_name": tname,
|
|
495
|
+
"transform_level": None,
|
|
496
|
+
"n_images": int(len(df_g)),
|
|
497
|
+
})
|
|
498
|
+
g2 = out_df.groupby(["transform_id","transform_name","transform_level"], dropna=False)
|
|
499
|
+
for (tid, tname, lvl), df_g in g2:
|
|
500
|
+
rows.append({
|
|
501
|
+
"scope": "attack_level",
|
|
502
|
+
"transform_id": int(tid) if pd.notna(tid) else None,
|
|
503
|
+
"transform_name": tname,
|
|
504
|
+
"transform_level": lvl,
|
|
505
|
+
"n_images": int(len(df_g)),
|
|
506
|
+
})
|
|
507
|
+
mdf = pd.DataFrame(rows)
|
|
508
|
+
mp = out["metrics_output"]["path"]
|
|
509
|
+
mc = out["metrics_output"].get("compression","snappy")
|
|
510
|
+
_ensure_dir(os.path.dirname(mp))
|
|
511
|
+
if out["metrics_output"].get("type","parquet") == "csv": mdf.to_csv(mp, index=False)
|
|
512
|
+
else: mdf.to_parquet(mp, index=False, compression=mc)
|
|
513
|
+
|
|
514
|
+
return out["table_output"]["path"]
|
|
515
|
+
|
|
516
|
+
if __name__ == "__main__":
|
|
517
|
+
import argparse
|
|
518
|
+
parser = argparse.ArgumentParser()
|
|
519
|
+
parser.add_argument("--config", required=True, help="Path to adversarial_pipeline.yaml")
|
|
520
|
+
args = parser.parse_args()
|
|
521
|
+
with open(args.config, "r") as f:
|
|
522
|
+
cfg = yaml.safe_load(f)
|
|
523
|
+
out_path = run_pipeline(cfg)
|
|
524
|
+
print("Wrote:", out_path)
|