MCPcopy Create free account
hub / github.com/ChenWu98/cycle-diffusion / evaluate

Method evaluate

evaluation/multi_task.py:12–71  ·  view source on GitHub ↗
(self, images, model, weighted_loss, losses, dataset, split)

Source from the content-addressed store, hash-verified

10 self.meta_args = meta_args
11
12 def evaluate(self, images, model, weighted_loss, losses, dataset, split):
13 assert split in ['eval', 'test']
14 assert len(weighted_loss) == len(dataset) == len(dataset.data)
15 num_examples = len(dataset)
16 assert all(len(v) == num_examples for k, v in losses.items())
17 if isinstance(images, torch.Tensor):
18 assert images.shape[0] == num_examples
19 elif isinstance(images, (list, tuple)):
20 assert (
21 all(_images.shape[0] == num_examples for _images in images)
22 or all(_images is None for _images in images)
23 )
24 elif images is None:
25 pass
26 else:
27 raise TypeError()
28
29 # Gather evaluation data for each task.
30 name2eval_kwargs = dict()
31 for i in range(num_examples):
32 name = dataset.data[i]['name']
33 if name not in name2eval_kwargs:
34 name2eval_kwargs[name] = {
35 "images": [],
36 "model": model,
37 "weighted_loss": [],
38 "losses": {
39 k: [] for k in losses.keys()
40 },
41 "data": [],
42 }
43 if isinstance(images, torch.Tensor):
44 name2eval_kwargs[name]['images'].append(images[i])
45 elif isinstance(images, (list, tuple)):
46 name2eval_kwargs[name]['images'].append(
47 tuple(_images[i] if _images is not None else None for _images in images)
48 )
49 elif images is None:
50 name2eval_kwargs[name]['images'].append(None)
51 else:
52 raise TypeError()
53
54 name2eval_kwargs[name]['weighted_loss'].append(weighted_loss[i])
55 for k, v in losses.items():
56 name2eval_kwargs[name]['losses'][k].append(v[i])
57 name2eval_kwargs[name]['data'].append(dataset.data[i])
58
59 # Evaluate each task.
60 summary = dict()
61 for name, eval_kwargs in name2eval_kwargs.items():
62 arg_path = getattr(self.meta_args.arg_paths, name)
63 args = get_config(arg_path)
64 evaluator = get_evaluator(args.evaluation.evaluator_program)(args, self.meta_args)
65 summary_tmp = evaluator.evaluate(**eval_kwargs, split=split)
66 for key, metric in summary_tmp.items():
67 summary[f'{name}/{key}'] = metric
68
69 if len(summary) > 0:

Callers

nothing calls this directly

Calls 2

get_configFunction · 0.90
get_evaluatorFunction · 0.90

Tested by

no test coverage detected