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,807 @@
1
+ from __future__ import annotations
2
+ import os
3
+ import io
4
+ import json
5
+ import uuid
6
+ import hashlib
7
+ import random
8
+ import pathlib
9
+ import inspect
10
+ import base64
11
+ import time
12
+ from typing import Any, Dict, List, Optional, Tuple
13
+ import math
14
+ import shutil
15
+
16
+ import numpy as np
17
+ import pandas as pd
18
+ from PIL import Image
19
+
20
+ import yaml
21
+ import albumentations as A
22
+ import cv2
23
+ import pyarrow as pa
24
+ import pyarrow.parquet as pq
25
+ from tqdm import tqdm
26
+
27
+
28
+
29
+ def _to_tensor_for_model(img: np.ndarray, resize):
30
+ """Convert numpy image to tensor for model inference"""
31
+ import torch
32
+ if resize:
33
+ img = cv2.resize(img, (resize[1], resize[0]), interpolation=cv2.INTER_LINEAR)
34
+ t = torch.from_numpy(img).float() / 255.0 # [0,1]
35
+ t = t.permute(2, 0, 1) # HWC->CHW
36
+ return t
37
+
38
+ def _normalize_tensor(t, mean, std):
39
+ """Normalize tensor"""
40
+ import torch
41
+ mean = torch.tensor(mean, device=t.device).view(3, 1, 1)
42
+ std = torch.tensor(std, device=t.device).view(3, 1, 1)
43
+ return (t - mean) / std
44
+
45
+ def _predict_class(model, img_arr: np.ndarray, device, mean, std, resize):
46
+ """Run inference on a single image"""
47
+ import torch
48
+ t = _to_tensor_for_model(img_arr, resize)
49
+ t = t.unsqueeze(0).to(device) # Add batch dimension
50
+ t = _normalize_tensor(t, mean, std)
51
+
52
+ with torch.no_grad():
53
+ logits = model(t)
54
+ pred = logits.argmax(1).item()
55
+ return pred
56
+
57
+
58
+ try:
59
+ from albumentations.core.utils import set_seed as _albumentations_set_seed
60
+ except Exception:
61
+ def _albumentations_set_seed(seed: int):
62
+ # fallback : on s'appuie sur random.seed et np.random.seed (si version inf albu à 1.3)
63
+ return
64
+ # contrôle randomisation
65
+ def _seed32(*parts, base_seed: int = 0) -> int:
66
+ h = hashlib.blake2b(digest_size=8)
67
+ for p in parts:
68
+ h.update(str(p).encode('utf-8'))
69
+ h.update(b'|')
70
+ return (int.from_bytes(h.digest(), 'little') ^ (base_seed & 0xFFFFFFFF)) & 0xFFFFFFFF
71
+
72
+ def _seed_all(seed: int):
73
+ random.seed(seed)
74
+ np.random.seed(seed)
75
+ _albumentations_set_seed(seed)
76
+
77
+
78
+ # exemple de transformation custom
79
+ class ColorTemperatureShift(A.ImageOnlyTransform):
80
+ """Chauffe/refroidit l'image en décalant R/B."""
81
+ def __init__(self, delta: int = 0, always_apply: bool = False, p: float = 1.0):
82
+ super().__init__(always_apply, p)
83
+ self.delta = int(delta)
84
+
85
+ def apply(self, img: np.ndarray, **params) -> np.ndarray:
86
+ out = img.astype(np.int16)
87
+ out[..., 0] = np.clip(out[..., 0] + self.delta, 0, 255) # R
88
+ out[..., 2] = np.clip(out[..., 2] - self.delta, 0, 255) # B
89
+ return out.astype(np.uint8)
90
+
91
+ CUSTOM_TRANSFORMS = {
92
+ "ColorTemperatureShift": ColorTemperatureShift,
93
+ }
94
+
95
+ # utils
96
+ def _ensure_dir(path: str):
97
+ if path:
98
+ pathlib.Path(path).mkdir(parents=True, exist_ok=True)
99
+
100
+ def _pil_from_bytes(b: bytes) -> Image.Image:
101
+ return Image.open(io.BytesIO(b)).convert("RGB")
102
+
103
+ def _np_from_pil(img: Image.Image) -> np.ndarray:
104
+ return np.array(img)
105
+
106
+ def _load_image_from_row(
107
+ row: pd.Series,
108
+ image_source: str,
109
+ col_image_bytes: Optional[str],
110
+ col_file_path: Optional[str],
111
+ base_dir: Optional[str] = None,
112
+ ) -> np.ndarray:
113
+ if image_source == "image_bytes":
114
+ assert col_image_bytes and col_image_bytes in row, f"Colonne '{col_image_bytes}' absente"
115
+ b = row[col_image_bytes]
116
+ if isinstance(b, str):
117
+ try:
118
+ b = b.encode("latin1")
119
+ except Exception:
120
+ try:
121
+ b = base64.b64decode(b)
122
+ except Exception as e:
123
+ raise ValueError("image_bytes: ni latin1 ni base64.") from e
124
+ img = _pil_from_bytes(b)
125
+ return _np_from_pil(img)
126
+
127
+ assert col_file_path and col_file_path in row, f"Colonne '{col_file_path}' absente"
128
+ fp = str(row[col_file_path])
129
+ if base_dir and not os.path.isabs(fp):
130
+ fp = os.path.join(base_dir, fp)
131
+ arr = cv2.imread(fp, cv2.IMREAD_UNCHANGED)
132
+ if arr is None:
133
+ raise FileNotFoundError(f"Impossible de lire l'image: {fp}")
134
+ if arr.ndim == 2:
135
+ arr = cv2.cvtColor(arr, cv2.COLOR_GRAY2RGB)
136
+ else:
137
+ arr = cv2.cvtColor(arr, cv2.COLOR_BGR2RGB)
138
+ return arr
139
+
140
+ def _encode_image(arr_rgb: np.ndarray, ext: str = "png") -> bytes:
141
+ if ext.lower() in ("jpg", "jpeg"):
142
+ ok, buf = cv2.imencode(".jpg", cv2.cvtColor(arr_rgb, cv2.COLOR_RGB2BGR), [int(cv2.IMWRITE_JPEG_QUALITY), 95])
143
+ elif ext.lower() in ("tif", "tiff"):
144
+ ok, buf = cv2.imencode(".tif", cv2.cvtColor(arr_rgb, cv2.COLOR_RGB2BGR))
145
+ else:
146
+ ok, buf = cv2.imencode(".png", cv2.cvtColor(arr_rgb, cv2.COLOR_RGB2BGR))
147
+ if not ok:
148
+ raise RuntimeError("Échec encodage image")
149
+ return buf.tobytes()
150
+
151
+ # utils paramètres pour les transformations
152
+ def _resolve_transform_class(name: str):
153
+ if name in CUSTOM_TRANSFORMS:
154
+ return CUSTOM_TRANSFORMS[name]
155
+ return getattr(A, name)
156
+
157
+ def _coerce_pair_param(param_name: str, value):
158
+ if isinstance(param_name, str) and (param_name.endswith("_range") or param_name.endswith("_limit")):
159
+ if isinstance(value, (int, float)):
160
+ return (value, value)
161
+ if isinstance(value, list):
162
+ if len(value) == 2:
163
+ return (value[0], value[1])
164
+ if len(value) == 1:
165
+ return (value[0], value[0])
166
+ return value
167
+
168
+ #
169
+ def _normalize_params_by_version(name: str, params: Dict[str, Any], allowed: set) -> Dict[str, Any]:
170
+ p = dict(params)
171
+ p.pop("level", None)
172
+
173
+ if name == "RandomFog":
174
+ if "fog_coef_range" in allowed and ("fog_coef_lower" in p or "fog_coef_upper" in p):
175
+ low = p.pop("fog_coef_lower", p.get("fog_coef_upper", 0.3))
176
+ up = p.pop("fog_coef_upper", p.get("fog_coef_lower", 1.0))
177
+ p["fog_coef_range"] = (low, up)
178
+
179
+ if name == "RandomSunFlare":
180
+ if "angle_range" in allowed and ("angle_lower" in p or "angle_upper" in p):
181
+ p["angle_range"] = (p.pop("angle_lower", 0.0), p.pop("angle_upper", 1.0))
182
+ if "num_flare_circles_range" in allowed and ("num_flare_circles_lower" in p or "num_flare_circles_upper" in p):
183
+ p["num_flare_circles_range"] = (p.pop("num_flare_circles_lower", 6), p.pop("num_flare_circles_upper", 10))
184
+
185
+ if name == "RandomShadow":
186
+ if "num_shadows_limit" in allowed and ("num_shadows_lower" in p or "num_shadows_upper" in p):
187
+ low = p.pop("num_shadows_lower", 1)
188
+ up = p.pop("num_shadows_upper", max(1, low))
189
+ p["num_shadows_limit"] = (low, up)
190
+
191
+ if name == "RandomRain":
192
+ if "slant_range" in allowed and ("slant_lower" in p or "slant_upper" in p):
193
+ p["slant_range"] = (p.pop("slant_lower", -10), p.pop("slant_upper", 10))
194
+
195
+ if name == "RandomSnow":
196
+ if "snow_point_range" in allowed and ("snow_point_lower" in p or "snow_point_upper" in p):
197
+ p["snow_point_range"] = (p.pop("snow_point_lower", 0.1), p.pop("snow_point_upper", 0.3))
198
+
199
+ if name == "CoarseDropout":
200
+ if "num_holes_range" in allowed and ("min_holes" in p or "max_holes" in p):
201
+ lo = p.pop("min_holes", None); hi = p.pop("max_holes", None)
202
+ if lo is None and hi is not None: lo = hi
203
+ if hi is None and lo is not None: hi = lo
204
+ if lo is not None and hi is not None:
205
+ p["num_holes_range"] = (int(lo), int(hi))
206
+ if "hole_height_range" in allowed and ("min_height" in p or "max_height" in p):
207
+ lo = p.pop("min_height", None); hi = p.pop("max_height", None)
208
+ if lo is None and hi is not None: lo = hi
209
+ if hi is None and lo is not None: hi = lo
210
+ if lo is not None and hi is not None:
211
+ p["hole_height_range"] = (int(lo), int(hi))
212
+ if "hole_width_range" in allowed and ("min_width" in p or "max_width" in p):
213
+ lo = p.pop("min_width", None); hi = p.pop("max_width", None)
214
+ if lo is not None and hi is not None:
215
+ p["hole_width_range"] = (int(lo), int(hi))
216
+ if "fill" in allowed and "fill_value" in p and "fill" not in p:
217
+ p["fill"] = p.pop("fill_value")
218
+ if "fill_mask" in allowed and "mask_fill_value" in p and "fill_mask" not in p:
219
+ p["fill_mask"] = p.pop("mask_fill_value")
220
+
221
+ for k in list(p.keys()):
222
+ if isinstance(k, str) and (k.endswith("_range") or k.endswith("_limit")):
223
+ p[k] = _coerce_pair_param(k, p[k])
224
+ return p
225
+
226
+ def _sanitize_params(transform_cls, name: str, params: Dict[str, Any]) -> Tuple[Dict[str, Any], Dict[str, Any]]:
227
+ sig = inspect.signature(transform_cls.__init__)
228
+ allowed = set(sig.parameters.keys()) - {"self", "args", "kwargs"}
229
+ normalized = _normalize_params_by_version(name, dict(params), allowed)
230
+ clean = {k: v for k, v in normalized.items() if k in allowed}
231
+ ignored = {k: v for k, v in normalized.items() if k not in allowed}
232
+ return clean, ignored
233
+
234
+ def _build_transform_instance(name: str, params: Dict[str, Any]) -> Tuple[A.BasicTransform, Dict[str, Any], Dict[str, Any]]:
235
+ cls = _resolve_transform_class(name)
236
+ clean, ignored = _sanitize_params(cls, name, params)
237
+ inst = cls(**clean)
238
+ return inst, clean, ignored
239
+
240
+ # def _iter_variants(transform_cfg: Dict[str, Any]):
241
+ # # (compat historique — non utilisé quand on active le sampling par image/transform unique)
242
+ # name = transform_cfg["name"]
243
+ # base_params = dict(transform_cfg.get("params", {}))
244
+
245
+ # if "level_overrides" in transform_cfg:
246
+ # for overrides in transform_cfg["level_overrides"]:
247
+ # params = dict(base_params)
248
+ # level = overrides.get("level", None)
249
+ # for k, v in overrides.items():
250
+ # if k == "level": continue
251
+ # params[k] = _coerce_pair_param(k, v)
252
+ # instance, effective, ignored = _build_transform_instance(name, params)
253
+ # yield {"tname": name, "level": level, "params": effective, "ignored": ignored, "instance": instance}
254
+ # return
255
+
256
+ # levels = transform_cfg.get("levels")
257
+ # level_param = transform_cfg.get("level_param")
258
+ # if levels and level_param:
259
+ # for lvl in levels:
260
+ # params = dict(base_params)
261
+ # params[level_param] = _coerce_pair_param(level_param, lvl)
262
+ # instance, effective, ignored = _build_transform_instance(name, params)
263
+ # yield {"tname": name, "level": lvl, "params": effective, "ignored": ignored, "instance": instance}
264
+ # return
265
+
266
+ # samples_per_image = int(transform_cfg.get("samples_per_image", 0))
267
+ # if samples_per_image > 0:
268
+ # for k in range(samples_per_image):
269
+ # instance, effective, ignored = _build_transform_instance(name, base_params)
270
+ # yield {"tname": name, "level": k, "params": effective, "ignored": ignored, "instance": instance}
271
+ # return
272
+
273
+ # instance, effective, ignored = _build_transform_instance(name, base_params)
274
+ # yield {"tname": name, "level": None, "params": effective, "ignored": ignored, "instance": instance}
275
+
276
+ def _normalize_weights(ws: Optional[List[float]], n: int) -> Optional[List[float]]:
277
+ if not ws:
278
+ return None
279
+ if len(ws) != n:
280
+ return None
281
+ w = np.array(ws, dtype=float)
282
+ s = w.sum()
283
+ if s <= 0:
284
+ return None
285
+ return (w / s).tolist()
286
+
287
+ def _pick_transform_for_image(transforms_cfg: List[Dict[str, Any]], pipeline_cfg: Dict[str, Any]) -> Dict[str, Any]:
288
+ # poids au niveau pipeline (liste) ou par transform via 'select_weight' (si one_transform_per_image: false)
289
+ n = len(transforms_cfg)
290
+ weights = pipeline_cfg.get("transform_weights")
291
+ if not weights:
292
+ weights = [t.get("select_weight") for t in transforms_cfg] if any("select_weight" in t for t in transforms_cfg) else None
293
+ p = _normalize_weights(weights, n)
294
+ idx = int(np.random.choice(n, p=p))
295
+ return transforms_cfg[idx]
296
+
297
+ def _sample_uniform_level(t_cfg: Dict[str, Any]) -> Optional[float]:
298
+ sp = t_cfg.get("sampling")
299
+ if sp and sp.get("mode", "uniform") == "uniform":
300
+ lo, hi = sp.get("range", [0.0, 0.0])
301
+ lo, hi = float(lo), float(hi)
302
+ L = float(np.random.uniform(lo, hi))
303
+ rd = sp.get("round")
304
+ if rd is not None:
305
+ L = round(L, int(rd))
306
+ return L
307
+
308
+ # fallback: si pas de sampling → si 'levels' existe, on choisit 1 niveau au hasard
309
+ if "levels" in t_cfg and t_cfg["levels"]:
310
+ return float(np.random.choice(t_cfg["levels"]))
311
+
312
+ return None
313
+
314
+
315
+ class LevelMonitor: # monitoring de traitmenet
316
+ def __init__(self, every: int = 0, log_on_start: bool = True, log_on_finish: bool = True):
317
+ self.every = int(every)
318
+ self.log_on_start = bool(log_on_start)
319
+ self.log_on_finish = bool(log_on_finish)
320
+ self.state: Dict[tuple, Dict[str, Any]] = {}
321
+
322
+ def _fmt_key(self, tid, tname, level):
323
+ lvl = "none" if level is None else level
324
+ return f"id={tid} | name={tname} | level={lvl}"
325
+
326
+ def start_if_needed(self, tid, tname, level):
327
+ key = (tid, tname, level)
328
+ if key not in self.state:
329
+ self.state[key] = {"n": 0, "t0": time.time(), "sources": set()}
330
+ if self.log_on_start:
331
+ print(f"[start] {self._fmt_key(tid, tname, level)}")
332
+
333
+ def tick(self, tid, tname, level, parent_id: str):
334
+ key = (tid, tname, level)
335
+ s = self.state[key]
336
+ s["n"] += 1
337
+ if parent_id is not None:
338
+ s["sources"].add(parent_id)
339
+ if self.every > 0 and (s["n"] % self.every == 0):
340
+ elapsed = max(1e-6, time.time() - s["t0"])
341
+ rate = s["n"] / elapsed
342
+ print(f"[proc] {self._fmt_key(tid, tname, level)}: n={s['n']} | sources={len(s['sources'])} | {elapsed:.1f}s | ~{rate:.1f} img/s")
343
+
344
+ def finish_all(self):
345
+ if not self.log_on_finish:
346
+ return
347
+ print("\n=== résumé par type/niveau ===")
348
+ for (tid, tname, level), s in sorted(self.state.items(), key=lambda kv: (kv[0][0], str(kv[0][2]))):
349
+ elapsed = max(1e-6, time.time() - s["t0"])
350
+ rate = s["n"] / elapsed
351
+ print(f"[done] id={tid:<2} name={tname:<18} level={('none' if level is None else level):<6} "
352
+ f"n={s['n']:<6} sources={len(s['sources']):<5} time={elapsed:.1f}s rate={rate:.1f} img/s")
353
+
354
+ # runner
355
+ def _stable_uid(base_id: str, tname: str, level: Any, params: Dict[str, Any]) -> str:
356
+ h = hashlib.sha1()
357
+ payload = json.dumps({"base": base_id, "tname": tname, "level": level, "params": params}, sort_keys=True).encode("utf-8")
358
+ h.update(payload)
359
+ return h.hexdigest()[:10]
360
+
361
+ class SafeDict(dict):
362
+ def __missing__(self, key):
363
+ return ""
364
+
365
+ def _write_table_single(df: pd.DataFrame, table_out_path: str, out_type: str, compression: Optional[str]):
366
+ _ensure_dir(os.path.dirname(table_out_path))
367
+ if out_type == "csv":
368
+ df.to_csv(table_out_path, index=False)
369
+ elif out_type == "parquet":
370
+ df.to_parquet(table_out_path, index=False, compression=compression)
371
+ else:
372
+ raise ValueError(f"Unsupported table output type: {out_type}")
373
+
374
+ def _clear_partition_dir(root_path: str, partition_cols: List[str], key_values: List[str]):
375
+ path = root_path
376
+ for col, val in zip(partition_cols, key_values):
377
+ path = os.path.join(path, f"{col}={val}")
378
+ shutil.rmtree(path, ignore_errors=True)
379
+
380
+ def _write_parquet_dataset(
381
+ df: pd.DataFrame,
382
+ root_path: str,
383
+ partition_cols: List[str],
384
+ compression: Optional[str],
385
+ existing_data_behavior: str = "overwrite_or_ignore",
386
+ max_rows_per_file: Optional[int] = None,
387
+ files_per_partition_min: Optional[int] = None,
388
+ max_rows_per_group: Optional[int] = None,
389
+ ):
390
+ if df.empty:
391
+ _ensure_dir(root_path)
392
+ return
393
+
394
+ df2 = df.copy()
395
+ for col in partition_cols:
396
+ if col not in df2.columns:
397
+ raise ValueError(f"partition col '{col}' absente de la table")
398
+ if "transform_level" in partition_cols:
399
+ df2["transform_level"] = df2["transform_level"].apply(lambda v: "none" if pd.isna(v) else str(v))
400
+ for col in partition_cols:
401
+ if col != "transform_level":
402
+ df2[col] = df2[col].astype(str)
403
+
404
+ _ensure_dir(root_path)
405
+ mrf = int(max_rows_per_file) if max_rows_per_file is not None else None
406
+ mrg = int(max_rows_per_group) if max_rows_per_group is not None else None
407
+ if mrf is not None and (mrg is None or mrg > mrf):
408
+ mrg = mrf
409
+
410
+ if not files_per_partition_min or int(files_per_partition_min) <= 1:
411
+ table = pa.Table.from_pandas(df2, preserve_index=False)
412
+ kwargs = {}
413
+ if mrf is not None:
414
+ kwargs["max_rows_per_file"] = mrf
415
+ if mrg is not None:
416
+ kwargs["max_rows_per_group"] = mrg
417
+ kwargs["row_group_size"] = mrg
418
+ pq.write_to_dataset(
419
+ table,
420
+ root_path=root_path,
421
+ partition_cols=partition_cols,
422
+ compression=compression if compression else None,
423
+ existing_data_behavior=existing_data_behavior,
424
+ **kwargs,
425
+ )
426
+ return
427
+
428
+ grouped = df2.groupby(partition_cols, dropna=False, sort=False)
429
+ for key_vals, part in grouped:
430
+ if not isinstance(key_vals, tuple):
431
+ key_vals = (key_vals,)
432
+ key_vals_str = [str(v) for v in key_vals]
433
+
434
+ if existing_data_behavior == "delete_matching":
435
+ _clear_partition_dir(root_path, partition_cols, key_vals_str)
436
+
437
+ nrows = len(part)
438
+ by_size = math.ceil(nrows / mrf) if mrf else 1
439
+ n_files = max(int(files_per_partition_min), by_size)
440
+ rows_per_shard = max(1, math.ceil(nrows / n_files))
441
+
442
+ for i in range(0, nrows, rows_per_shard):
443
+ shard = part.iloc[i:i+rows_per_shard]
444
+ table = pa.Table.from_pandas(shard, preserve_index=False)
445
+ pq.write_to_dataset(
446
+ table,
447
+ root_path=root_path,
448
+ partition_cols=partition_cols,
449
+ compression=compression if compression else None,
450
+ existing_data_behavior="overwrite_or_ignore",
451
+ )
452
+
453
+ def run_pipeline(cfg: Dict[str, Any]) -> str:
454
+ ap = cfg["augment_pipeline"]
455
+
456
+ # [OPTIONNEL] Charger le modèle pour l'inférence
457
+ model = None
458
+ device = None
459
+ model_mean = None
460
+ model_std = None
461
+ model_resize = None
462
+ model_name = None # nom du modèle
463
+ use_inference = False
464
+
465
+ if "model" in ap:
466
+ import torch
467
+ import torchvision.models as models
468
+
469
+ mcfg = ap["model"]
470
+ model_resize = mcfg.get("resize")
471
+ model_mean = mcfg.get("mean", [0.0, 0.0, 0.0])
472
+ model_std = mcfg.get("std", [1.0, 1.0, 1.0])
473
+ device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
474
+ if device.type == "cuda":
475
+ try:
476
+ import torch as _torch
477
+ _ = _torch.tensor([1.0], device=device) + 1.0
478
+ _torch.cuda.synchronize()
479
+ except Exception as e:
480
+ print(f"CUDA non utilisable avec cette build ({str(e)[:100]}). Bascule sur CPU.")
481
+ device = torch.device("cpu")
482
+
483
+ ckpt = mcfg["ckpt_path"]
484
+ model_name = os.path.basename(ckpt) # nom du modèle
485
+
486
+ # Set quantization backend (needed for quantized models)
487
+ if device.type == 'cpu':
488
+ supported_engines = torch.backends.quantized.supported_engines
489
+ if 'qnnpack' in supported_engines:
490
+ torch.backends.quantized.engine = 'qnnpack'
491
+ elif 'fbgemm' in supported_engines:
492
+ torch.backends.quantized.engine = 'fbgemm'
493
+ print(f"🔧 Quantization engine: {torch.backends.quantized.engine}")
494
+
495
+ # Try to load as TorchScript first (for .pt files), otherwise as regular checkpoint
496
+ try:
497
+ # Try TorchScript load
498
+ model = torch.jit.load(ckpt, map_location=device)
499
+ print(f"✅ Modèle TorchScript chargé : {model_name}")
500
+ except RuntimeError as e:
501
+ # Not a TorchScript file, try regular checkpoint
502
+ checkpoint = torch.load(ckpt, map_location=device, weights_only=False)
503
+
504
+ # Handle different checkpoint formats
505
+ if isinstance(checkpoint, dict):
506
+ if "model_state_dict" in checkpoint:
507
+ # Standard checkpoint format with model state dict
508
+ arch = checkpoint.get("arch", "resnet18")
509
+ num_classes = checkpoint.get("num_classes", 1000)
510
+
511
+
512
+ if arch == "resnet18":
513
+ model = models.resnet18(pretrained=False)
514
+ model.fc = torch.nn.Linear(model.fc.in_features, num_classes)
515
+ else:
516
+ raise ValueError(f"Unsupported architecture: {arch}")
517
+
518
+ # Load state dict
519
+ model.load_state_dict(checkpoint["model_state_dict"])
520
+ else:
521
+ # Assume the dict itself is the model
522
+ model = checkpoint
523
+ else:
524
+ # Direct model object
525
+ model = checkpoint
526
+ print(f"Modèle checkpoint chargé : {model_name}")
527
+
528
+ if hasattr(model, "eval"): model.eval()
529
+ model = model.to(device)
530
+ use_inference = True
531
+ print(f"🔬 Modèle chargé pour l'inférence sur {device}")
532
+ else:
533
+ print(" Mode génération pure (pas d'inférence)")
534
+
535
+ # loading des données
536
+ dl = ap["dataloader"]
537
+ input_type = dl["type"]
538
+ input_path = dl["path"]
539
+ batch_size = int(dl.get("batch_size", 10_000))
540
+ columns = dl["columns"]
541
+ image_source = columns.get("image_source", "image_bytes")
542
+ col_bytes = columns.get("image_bytes")
543
+ col_path = columns.get("file_path")
544
+ base_dir = dl.get("base_dir")
545
+
546
+ # outputs - parquet et partionning
547
+ outputs = ap["outputs"]
548
+ img_out = outputs["image_output"]
549
+ file_ext = img_out.get("file_ext", "png")
550
+
551
+ table_out = outputs["table_output"]
552
+ table_out_type = table_out["type"]
553
+ table_out_path = table_out["path"]
554
+ compression = table_out.get("compression", None)
555
+ # Only enable dataset mode if BOTH dataset=true AND partition_by is provided
556
+ dataset_mode = bool(table_out.get("dataset", False)) and bool(table_out.get("partition_by"))
557
+ partition_cols = list(table_out.get("partition_by", []))
558
+ existing_behavior = table_out.get("existing_data_behavior", "overwrite_or_ignore")
559
+ max_rows_per_file = table_out.get("max_rows_per_file")
560
+ files_per_partition_min = table_out.get("files_per_partition_min")
561
+ max_rows_per_group = table_out.get("max_rows_per_group")
562
+
563
+ # metrics
564
+ metrics_out = outputs.get("metrics_output", None)
565
+ ordered_front = outputs["columns"]["ordered_front"]
566
+ include_src = outputs["columns"].get("include_source_columns", False)
567
+
568
+ # run
569
+ run_opts = ap.get("run", {})
570
+ seed = int(run_opts.get("seed", 123))
571
+ deterministic_ids = bool(run_opts.get("deterministic_ids", False))
572
+ id_pattern = run_opts.get("id_pattern", "{base}__{tname}__{level}__{uid}")
573
+ output_sample_type = run_opts.get("output_sample_type", "perturbation")
574
+
575
+ # monitor
576
+ monitor_cfg = ap.get("monitor", {})
577
+ per_level_every = int(monitor_cfg.get("log_every_per_level", 0))
578
+ log_on_start = bool(monitor_cfg.get("log_on_start", True))
579
+ log_on_finish = bool(monitor_cfg.get("log_on_finish", True))
580
+ mon = LevelMonitor(per_level_every, log_on_start, log_on_finish)
581
+
582
+ _seed_all(seed)
583
+
584
+ # read
585
+ if input_type == "csv":
586
+ full_df = pd.read_csv(input_path)
587
+ elif input_type == "parquet":
588
+ full_df = pd.read_parquet(input_path)
589
+ else:
590
+ raise ValueError("'")
591
+
592
+ pipelines = ap.get("augmentation_pipelines", [])
593
+ if not pipelines:
594
+ raise ValueError("aucune pipeline/transforms fournie")
595
+
596
+ all_records: List[Dict[str, Any]] = []
597
+ n = len(full_df)
598
+ loaded_ok, loaded_fail = 0, 0
599
+
600
+ num_pipelines = len(pipelines)
601
+ num_batches = (n + batch_size - 1) // batch_size
602
+
603
+ print(f"\n{'='*60}")
604
+ print(f"Démarrage du pipeline d'augmentation")
605
+ print(f"{n} images | {num_pipelines} pipeline(s) | {num_batches} batches")
606
+ print(f"{'='*60}\n")
607
+
608
+ pbar = tqdm(total=n, desc="Traitement images", unit="img", ncols=100)
609
+
610
+ for start in range(0, n, batch_size):
611
+ end = min(start + batch_size, n)
612
+ batch = full_df.iloc[start:end]
613
+
614
+ for idx, (_, row) in enumerate(batch.iterrows()):
615
+ try:
616
+ arr = _load_image_from_row(row, image_source, col_bytes, col_path, base_dir=base_dir)
617
+ loaded_ok += 1
618
+ except Exception:
619
+ loaded_fail += 1
620
+ continue
621
+
622
+ base_id = row.get(columns.get("sample_id", "sample_id"), f"row{start}")
623
+ class_id = row.get(columns.get("class_id", "class_id"), None)
624
+ class_name = row.get(columns.get("class_name", "class_name"), None)
625
+ split = row.get(columns.get("split", "split"), None)
626
+
627
+ for pipe_idx, pipe_cfg in enumerate(pipelines):
628
+ # on a la décision apply_prob, reseedée par image + pipeline
629
+ apply_prob = float(pipe_cfg.get("apply_prob", 1.0))
630
+ if apply_prob < 1.0:
631
+ _seed_all(_seed32(base_id, pipe_idx, "apply_prob", base_seed=seed))
632
+ if random.random() > apply_prob:
633
+ continue # skip de manière déterministe
634
+
635
+ # option de resizing
636
+ resize = pipe_cfg.get("resize")
637
+ if resize:
638
+ H, W = resize
639
+ arr_in = A.Resize(height=H, width=W)(image=arr)["image"]
640
+ else:
641
+ arr_in = arr
642
+
643
+ transforms_cfg = pipe_cfg.get("transforms", [])
644
+ if not transforms_cfg:
645
+ continue
646
+
647
+ # Choix DU transform (1 seul) de façon déterministe
648
+ one_transform = bool(pipe_cfg.get("one_transform_per_image", True))
649
+ if one_transform:
650
+ _seed_all(_seed32(base_id, pipe_idx, "choose_transform", base_seed=seed))
651
+ chosen_t = _pick_transform_for_image(transforms_cfg, pipe_cfg)
652
+ t_cfgs = [chosen_t]
653
+ else:
654
+ t_cfgs = transforms_cfg # compat: plusieurs sorties (ancienne version)
655
+
656
+ for t_cfg in t_cfgs:
657
+ transform_alias = t_cfg.get("alias")
658
+ transform_id_value = int(t_cfg.get("id", 0))
659
+ alb_name = t_cfg["name"]
660
+ level_param = t_cfg.get("level_param")
661
+
662
+ # tirage du niveau: reseed dédié pour le level
663
+ _seed_all(_seed32(base_id, transform_id_value, "level", base_seed=seed))
664
+ level = _sample_uniform_level(t_cfg)
665
+
666
+ params = dict(t_cfg.get("params", {}))
667
+ if level_param is not None and level is not None:
668
+ params[level_param] = _coerce_pair_param(level_param, level)
669
+
670
+ t, params_used, _ignored = _build_transform_instance(alb_name, params)
671
+ tname_display = transform_alias or alb_name
672
+
673
+ # barre de progression avec la transformation en cours
674
+ level_str = f"lvl={level:.3f}" if level is not None else "lvl=N/A"
675
+ pbar.set_description(f"🎨 {tname_display} ({level_str})")
676
+
677
+ _seed_all(_seed32(base_id, transform_id_value, level, "apply", base_seed=seed))
678
+
679
+ mon.start_if_needed(transform_id_value, tname_display, level)
680
+ out_arr = A.Compose([t])(image=arr_in)["image"]
681
+ img_bytes = _encode_image(out_arr, file_ext)
682
+
683
+ # Calculate hash of the transformed image
684
+ img_hash = hashlib.sha256(img_bytes).hexdigest()
685
+
686
+ # Generate UUIDs
687
+ sample_uuid_val = str(uuid.uuid4())
688
+ parent_sample_uuid_val = row.get(columns.get("sample_uuid", "sample_uuid"), str(uuid.uuid4()))
689
+
690
+ # iD stable optionnel
691
+ uid = _stable_uid(base_id, tname_display, level, params_used) if deterministic_ids else str(uuid.uuid4())[:8]
692
+ fmt_vars = SafeDict({
693
+ "base": base_id,
694
+ "tname": tname_display,
695
+ "level": ("none" if level is None else str(level)),
696
+ "uid": uid
697
+ })
698
+ child_id = id_pattern.format_map(fmt_vars)
699
+
700
+ rec = {
701
+ "sample_uuid": sample_uuid_val,
702
+ "parent_sample_uuid": parent_sample_uuid_val,
703
+ "sample_id": base_id,
704
+ "sample_type": output_sample_type,
705
+ "class_id": class_id,
706
+ "class_name": class_name,
707
+ "split": split,
708
+ "transform_id": transform_id_value,
709
+ "transform_name": tname_display,
710
+ "transform_level": level,
711
+ "transform_params": json.dumps(params_used, ensure_ascii=False),
712
+ "image_bytes": img_bytes,
713
+ "image_hash": img_hash,
714
+ }
715
+
716
+ if include_src:
717
+ for c in full_df.columns:
718
+ if c not in rec:
719
+ rec[f"src__{c}"] = row[c]
720
+
721
+ all_records.append(rec)
722
+ mon.tick(transform_id_value, tname_display, level, base_id)
723
+
724
+ # Mise à jour de la barre de progression après chaque image
725
+ pbar.update(1)
726
+
727
+ pbar.close()
728
+
729
+ print(f"\n{'='*60}")
730
+ print(f"pipeline terminé")
731
+ print(f"{'='*60}\n")
732
+
733
+ out_df = pd.DataFrame(all_records)
734
+
735
+ if not out_df.empty:
736
+ # Filter ordered_front to only include columns that actually exist in the DataFrame
737
+ front = [c for c in ordered_front if c in out_df.columns]
738
+ rest = [c for c in out_df.columns if c not in front]
739
+ out_df = out_df[front + rest]
740
+
741
+ print(f"Sources chargées: {loaded_ok}, échecs: {loaded_fail}")
742
+ print(f"Augmentations générées: {len(out_df)}")
743
+
744
+ if dataset_mode:
745
+ _write_parquet_dataset(
746
+ out_df,
747
+ table_out_path,
748
+ partition_cols or ["transform_name", "transform_level"],
749
+ compression,
750
+ existing_data_behavior=existing_behavior,
751
+ max_rows_per_file=int(max_rows_per_file) if max_rows_per_file else None,
752
+ files_per_partition_min=int(files_per_partition_min) if files_per_partition_min else None,
753
+ max_rows_per_group=int(max_rows_per_group) if max_rows_per_group else None,
754
+ )
755
+ print(f" dataset Parquet partitionné -> {table_out_path}")
756
+ else:
757
+ _write_table_single(out_df, table_out_path, table_out_type, compression)
758
+ print(f" {table_out_type} écrit -> {table_out_path}")
759
+
760
+ if outputs.get("metrics_output"):
761
+ metrics = outputs["metrics_output"]
762
+ m_type = metrics.get("type", "parquet")
763
+ m_path = metrics["path"]
764
+ m_comp = metrics.get("compression")
765
+
766
+ rows: List[Dict[str, Any]] = []
767
+ if not out_df.empty:
768
+ g = out_df.groupby(["transform_id", "transform_name"], dropna=False)
769
+ for (tid, tname), df_g in g:
770
+ rows.append({
771
+ "scope": "transform",
772
+ "transform_id": int(tid) if pd.notna(tid) else None,
773
+ "transform_name": tname,
774
+ "transform_level": None,
775
+ "n_images": int(len(df_g)),
776
+ })
777
+ g2 = out_df.groupby(["transform_id", "transform_name", "transform_level"], dropna=False)
778
+ for (tid, tname, lvl), df_g in g2:
779
+ rows.append({
780
+ "scope": "transform_level",
781
+ "transform_id": int(tid) if pd.notna(tid) else None,
782
+ "transform_name": tname,
783
+ "transform_level": lvl,
784
+ "n_images": int(len(df_g)),
785
+ })
786
+ mdf = pd.DataFrame(rows)
787
+
788
+ _ensure_dir(os.path.dirname(m_path))
789
+
790
+ if m_type == "csv":
791
+ mdf.to_csv(m_path, index=False)
792
+ else:
793
+ mdf.to_parquet(m_path, index=False, compression=m_comp)
794
+ print(f" métriques écrites -> {m_path}")
795
+
796
+ mon.finish_all()
797
+ return table_out_path
798
+
799
+ if __name__ == "__main__":
800
+ import argparse
801
+
802
+ parser = argparse.ArgumentParser()
803
+ parser.add_argument("--config", required=True, help="Path to augment_pipeline.yaml")
804
+ args = parser.parse_args()
805
+ with open(args.config, "r") as f:
806
+ cfg = yaml.safe_load(f)
807
+ out_path = run_pipeline(cfg)