-
Notifications
You must be signed in to change notification settings - Fork 1
Expand file tree
/
Copy pathtest.py
More file actions
80 lines (66 loc) · 2.65 KB
/
Copy pathtest.py
File metadata and controls
80 lines (66 loc) · 2.65 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
# test.py
import logging
import os
import lightning as L
from lightning.pytorch.loggers import CSVLogger, WandbLogger
from deep_stylometry.callbacks import TestEvalCallback
from deep_stylometry.modules import DeepStylometry
from deep_stylometry.utils import train_utils
from deep_stylometry.utils.argparsers import TestArgparse
from deep_stylometry.utils.configs import BaseConfig
from deep_stylometry.utils.helpers import resolve_lightning_precision, set_seed
from deep_stylometry.utils.logger import logging_config
os.environ["PYTHONUNBUFFERED"] = "1"
os.environ["TOKENIZERS_PARALLELISM"] = "false"
logger = logging.getLogger(__name__)
set_seed()
logging_config()
if __name__ == "__main__":
args = TestArgparse.parse_known_args()
cfg = BaseConfig(mode="test").from_yaml(args.config_path)
# CLI overrides (allow per-job customisation without editing YAML)
if args.test_subset is not None:
cfg.data.test_subset = args.test_subset
if args.ds_name is not None:
cfg.data.ds_name = args.ds_name
logger.info("Preparing data module...")
dm = train_utils.setup_datamodule(
cfg=cfg,
processed_ds_dir=args.processed_ds_dir,
num_proc=args.num_proc,
cache_dir=args.cache_dir,
)
if cfg.model.pooling_method == "pli":
patch_tag = f"pli-{cfg.model.patch_method}-n{cfg.model.patch_size}"
else:
patch_tag = cfg.model.pooling_method
name = (
f"test__{cfg.model.base_checkpoint}__{cfg.data.ds_name}"
f"__pooling-{patch_tag}__{cfg.data.test_subset}"
f"__skip_list-{cfg.model.skip_list}"
).replace("/", "-").lower()
loggers = []
if cfg.test.use_wandb:
wandb_logger = WandbLogger(
project=cfg.project_name,
name=name,
)
loggers.append(wandb_logger)
# Add CSV logger by default
csv_logger = CSVLogger(save_dir=args.logs_dir, name=name)
loggers.append(csv_logger)
logger.info(f"Loading model from {args.checkpoint_path}...")
# Load weights from your saved checkpoint
model = DeepStylometry.load_from_checkpoint(args.checkpoint_path, cfg=cfg)
precision, _ = resolve_lightning_precision(cfg.test.precision)
# Force 1 GPU for testing to avoid distributed gather complications
trainer = L.Trainer(
accelerator=cfg.test.device,
devices=cfg.test.num_devices,
logger=loggers, # Set to your WandbLogger/CSVLogger if you want to save the test metrics remotely
callbacks=[TestEvalCallback(cfg)],
precision=precision,
)
logger.info("=== Starting Evaluation ===")
trainer.test(model=model, datamodule=dm)
logger.info("=== Evaluation Finished ===")