-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy patheval_main.py
More file actions
129 lines (105 loc) · 3.78 KB
/
Copy patheval_main.py
File metadata and controls
129 lines (105 loc) · 3.78 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
"""
Evaluation entry point for CamFlow model.
Following KISS, YAGNI and SOLID principles.
"""
import os
import torch
import logging
import argparse
from datetime import datetime
from accelerate import Accelerator
from accelerate.utils import DistributedDataParallelKwargs
from common.utils import Params, set_logger
from common.manager import Manager
from dataset.data_loader import fetch_dataloader
from model.net import fetch_net
from evaluators import GHOFEvaluator, IQAEvaluator
def parse_args():
"""Parse command line arguments."""
parser = argparse.ArgumentParser(description='Evaluate CamFlow model')
parser.add_argument('--model_dir',
default='experiments/CAHomo/',
help="Directory containing params.json")
parser.add_argument('--restore_file',
default='experiments/CAHomo/HEM.pth',
help="Path to model weights")
parser.add_argument('--only_weights',
action='store_true',
default=True,
help='Only use weights to load or load all train status')
parser.add_argument('--seed',
type=int,
default=230,
help='Random seed')
parser.add_argument('--enable_iqa',
action='store_true',
default=False,
help='Enable IQA metrics evaluation')
return parser.parse_args()
def setup_environment(args):
"""Setup evaluation environment."""
# Load params
json_path = os.path.join(args.model_dir, 'params.json')
assert os.path.isfile(json_path), f"No json configuration file found at {json_path}"
params = Params(json_path)
params.update(vars(args))
# Set CUDA
params.cuda = torch.cuda.is_available()
# Set random seeds
torch.manual_seed(args.seed)
if params.cuda:
torch.cuda.manual_seed(args.seed)
# Setup logging
logger = set_logger(os.path.join(args.model_dir, 'evaluate.log'))
return params, logger
def build_model_and_manager(params, logger):
"""Build model and evaluation manager."""
# Create dataloaders
dataloaders = fetch_dataloader(params)
# Create model
model = fetch_net(params)
# Setup DDP
kwargs = DistributedDataParallelKwargs(find_unused_parameters=True)
accelerator = Accelerator(split_batches=True, kwargs_handlers=[kwargs])
# Prepare model and data
model = accelerator.prepare(model)
for split in ['train', 'val', 'test', 'ghof']:
if split in dataloaders:
dataloaders[split] = accelerator.prepare(dataloaders[split])
# Create manager
manager = Manager(
model=model,
optimizer=None,
scheduler=None,
params=params,
dataloaders=dataloaders,
writer=None,
logger=logger,
accelerator=accelerator
)
# Load checkpoints
manager.load_checkpoints()
return manager
def main():
"""Main evaluation function."""
# Parse arguments and setup
args = parse_args()
params, logger = setup_environment(args)
# Build model and manager
manager = build_model_and_manager(params, logger)
# Create evaluators
evaluators = []
evaluators.append(GHOFEvaluator(manager))
if args.enable_iqa:
evaluators.append(IQAEvaluator(manager))
# Store evaluators in manager for sharing results
manager.evaluators = evaluators
# Run evaluation
if manager.accelerator.is_main_process:
logger.info("Starting evaluation")
for evaluator in evaluators:
evaluator.evaluate()
if manager.accelerator.is_main_process:
logger.info("Evaluation complete!")
if __name__ == '__main__':
main()