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.
@@ -0,0 +1,9 @@
1
+ __description__ = "Generating samples......"
2
+
3
+ from .augment_run import run_pipeline as RunAugmentPipeline
4
+ from .adversarial_run import run_pipeline as RunAdversarialPipeline
5
+
6
+ __all__ = [
7
+ "RunAugmentPipeline",
8
+ "RunAdversarialPipeline"
9
+ ]
@@ -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)