-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathmain.py
More file actions
130 lines (100 loc) · 3.66 KB
/
Copy pathmain.py
File metadata and controls
130 lines (100 loc) · 3.66 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
"""
=====================================================================
XAI Attack and Defense Framework - Pipeline Orchestrator
Author : Mithin Sagar S (https://github.com/mithinsagar)
=====================================================================
Runs the end-to-end pipeline (or individual phases) via a single CLI
entrypoint. Each phase delegates to a dedicated module under `training/`,
`attacks/`, `defenses/`, and `explainers/`.
Usage
-----
python main.py --phase all
python main.py --phase preprocess
python main.py --phase train_base
python main.py --phase explain
python main.py --phase attack
python main.py --phase fewshot
python main.py --phase defense
"""
from __future__ import annotations
import argparse
import sys
from typing import Callable, Dict
import config
from utils.logger import get_logger
from utils.seed import set_seed
LOGGER = get_logger("main")
# ---------------------------------------------------------------------
# Phase entrypoints (lazy imports to keep startup fast)
# ---------------------------------------------------------------------
def _phase_preprocess() -> None:
from data.preprocessor import preprocess_all_datasets
LOGGER.info("Phase 1: Preprocessing all datasets")
preprocess_all_datasets()
def _phase_train_base() -> None:
from training.train_base_models import train_all_base_models
LOGGER.info("Phase 2: Training all base models")
train_all_base_models()
def _phase_explain() -> None:
from explainers.explainer_utils import generate_all_explanations
LOGGER.info("Phase 3: Generating explanations")
generate_all_explanations()
def _phase_attack() -> None:
from attacks.attack_runner import run_all_attacks
LOGGER.info("Phase 4: Running attacks")
run_all_attacks()
def _phase_fewshot() -> None:
from training.train_fewshot import train_and_evaluate_fewshot
LOGGER.info("Phase 5: Few-shot vulnerability analysis")
train_and_evaluate_fewshot()
def _phase_defense() -> None:
from training.train_defense_models import train_all_defenses
LOGGER.info("Phase 6: Training defense models")
train_all_defenses()
PHASES: Dict[str, Callable[[], None]] = {
"preprocess": _phase_preprocess,
"train_base": _phase_train_base,
"explain": _phase_explain,
"attack": _phase_attack,
"fewshot": _phase_fewshot,
"defense": _phase_defense,
}
# ---------------------------------------------------------------------
# CLI
# ---------------------------------------------------------------------
def build_argparser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser(
description=(
"XAI Attack and Defense Framework - Pipeline Orchestrator. "
"Author: Mithin Sagar S "
"(https://github.com/mithinsagar)."
)
)
parser.add_argument(
"--phase",
choices=list(PHASES.keys()) + ["all"],
default="all",
help="Which pipeline phase to execute. Default: all.",
)
parser.add_argument(
"--seed",
type=int,
default=config.RANDOM_SEED,
help="Random seed used across NumPy, Python, PyTorch.",
)
return parser
def main(argv: list[str] | None = None) -> int:
args = build_argparser().parse_args(argv)
config.ensure_directories()
set_seed(args.seed)
LOGGER.info("Random seed set to %d", args.seed)
if args.phase == "all":
for name, phase_fn in PHASES.items():
LOGGER.info("--- Running phase: %s ---", name)
phase_fn()
else:
PHASES[args.phase]()
LOGGER.info("Pipeline complete.")
return 0
if __name__ == "__main__":
sys.exit(main())