craceplot 0.1.0__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.
Files changed (61) hide show
  1. craceplot/__init__.py +60 -0
  2. craceplot/_containers/__init__.py +0 -0
  3. craceplot/_containers/_core.py +1602 -0
  4. craceplot/_containers/_draw.py +104 -0
  5. craceplot/_inst/examples/acotsp/crace-3240078.stdout +838 -0
  6. craceplot/_inst/examples/acotsp/crace.stderr +8 -0
  7. craceplot/_inst/examples/acotsp/crace.test +216 -0
  8. craceplot/_inst/examples/acotsp/crace.train +622 -0
  9. craceplot/_inst/examples/cats200/crace-2814581.stdout +1665 -0
  10. craceplot/_inst/examples/cats200/crace.stderr +10 -0
  11. craceplot/_inst/examples/cats200/crace.test +1010 -0
  12. craceplot/_inst/examples/cats200/crace.train +655 -0
  13. craceplot/_plots/__init__.py +8 -0
  14. craceplot/_plots/_core.py +80 -0
  15. craceplot/_plots/parameters.py +1805 -0
  16. craceplot/_plots/quality.py +1406 -0
  17. craceplot/_scripts/__init__.py +4 -0
  18. craceplot/_scripts/_main.py +132 -0
  19. craceplot/_scripts/_utils.py +66 -0
  20. craceplot/_scripts/main +3 -0
  21. craceplot/_scripts/open_guide +12 -0
  22. craceplot/_settings/_description.py +168 -0
  23. craceplot/_settings/_options.json +299 -0
  24. craceplot/_utils/__init__.py +0 -0
  25. craceplot/_utils/_base.py +196 -0
  26. craceplot/_utils/_const.py +8 -0
  27. craceplot/_utils/_crace.py +42 -0
  28. craceplot/_utils/_format.py +191 -0
  29. craceplot/_vergit.py +2 -0
  30. craceplot/_version.py +24 -0
  31. craceplot/_vignettes/guide.ipynb +18943 -0
  32. craceplot-0.1.0.dist-info/METADATA +369 -0
  33. craceplot-0.1.0.dist-info/RECORD +61 -0
  34. craceplot-0.1.0.dist-info/WHEEL +5 -0
  35. craceplot-0.1.0.dist-info/entry_points.txt +2 -0
  36. craceplot-0.1.0.dist-info/licenses/LICENSE.md +21 -0
  37. craceplot-0.1.0.dist-info/top_level.txt +2 -0
  38. docs/_config.yml +51 -0
  39. docs/_static/css/custom.css +168 -0
  40. docs/_static/js/navbar.js +40 -0
  41. docs/_toc.yml +5 -0
  42. docs/index.md +4 -0
  43. docs/references/404.md +14 -0
  44. docs/references/authors.md +46 -0
  45. docs/references/citiation.bib +17 -0
  46. docs/references/functions/index.md +5 -0
  47. docs/references/functions/param_boxplot.ipynb +19 -0
  48. docs/references/index.md +44 -0
  49. docs/references/license.md +21 -0
  50. docs/references/others/favicon.ico +0 -0
  51. docs/references/others/favicon_io/about.txt +6 -0
  52. docs/references/others/favicon_io/android-chrome-192x192.png +0 -0
  53. docs/references/others/favicon_io/android-chrome-512x512.png +0 -0
  54. docs/references/others/favicon_io/apple-touch-icon.png +0 -0
  55. docs/references/others/favicon_io/favicon-16x16.png +0 -0
  56. docs/references/others/favicon_io/favicon-32x32.png +0 -0
  57. docs/references/others/favicon_io/site.webmanifest +1 -0
  58. docs/references/others/logo.png +0 -0
  59. docs/references/others/logo1.png +0 -0
  60. docs/references/plots/acotsp-parallelcoord-all.png +0 -0
  61. docs/requirements.txt +5 -0
@@ -0,0 +1,1406 @@
1
+ # craceplot._plot.parameters.py
2
+
3
+ from __future__ import annotations
4
+
5
+ import os
6
+ import sys
7
+ import re
8
+ import math
9
+ import inspect
10
+ import traceback
11
+
12
+ import numpy as np
13
+ import pandas as pd
14
+ import seaborn as sns
15
+ import matplotlib.pyplot as plt
16
+ import matplotlib.lines as mlines
17
+
18
+
19
+ from types import SimpleNamespace
20
+ from sklearn.utils import resample
21
+ from matplotlib.transforms import Bbox
22
+ from typing import TYPE_CHECKING, Union, Literal, Callable, TypeVar, ParamSpec
23
+
24
+
25
+ from craceplot._plots._core import *
26
+ from craceplot._utils._base import *
27
+ from craceplot._containers._core import bold, underline, reset, CplotOptions, Errors as CE
28
+
29
+
30
+ P = ParamSpec("P")
31
+ R = TypeVar("R")
32
+
33
+ if TYPE_CHECKING:
34
+ from crace.containers.crace_results import CraceResults
35
+
36
+
37
+ __all__ = []
38
+
39
+ def export(func: Callable[P, R]) -> Callable[P, R]:
40
+ """append func name to list __all__"""
41
+ __all__.append(func.__name__)
42
+ return func
43
+
44
+ args = SimpleNamespace(
45
+ # required parameters
46
+ data=None,
47
+ options=None,
48
+ source=None,
49
+ num=None,
50
+ select=None,
51
+ slice=True,
52
+ out_dir=None,
53
+ file_name=None,
54
+ title=None,
55
+ _console=True,
56
+ # optional parameters
57
+ showfliers=False,
58
+ showmeans=False,
59
+ stest=False,
60
+ dpi=800,
61
+ )
62
+
63
+ def add_slice(data: pd.DataFrame, slice: pd.DataFrame):
64
+ """add slice information to selected data"""
65
+
66
+ end = slice['end_time'].values
67
+ idx = np.searchsorted(end, data["end_time"].values, side="right")
68
+ new = data.assign(slice_finished=idx + 1)
69
+
70
+ return new
71
+
72
+ @enforce_types
73
+ def _check_options(name: str, args: SimpleNamespace):
74
+ """
75
+ check options
76
+
77
+ options provided by args have higher priority
78
+
79
+ each option will be read from args when plotting
80
+
81
+ ONLY modify args
82
+
83
+ # required parameters
84
+ data: CraceResults=None,
85
+ options: CplotOptions=None,
86
+ select: list=None,
87
+ slice: bool=True,
88
+ out_dir: str=None,
89
+ file_name: str=None,
90
+ title: str=None,
91
+ num: Literal['all', 'final', 'elites', None] = None,
92
+ source: Literal['training', 'test', 'testing', None] = None,
93
+ _console: bool=True,
94
+ # optional parameters
95
+ showfliers: bool=False,
96
+ showmeans: bool=False,
97
+ stest: bool=False,
98
+ dpi: int=800,
99
+ """
100
+ # data must be provided
101
+ if not args.data:
102
+ raise CE.OptionError(f"{bold}data{reset} must be provided for plotting.")
103
+
104
+ # check onlytest
105
+ onlytest = args.data.options.onlytest.value
106
+ if onlytest:
107
+ print(f"# The provided crace log files only has results for test part")
108
+ args.source = 'testing'
109
+ args.num = 'final'
110
+
111
+ # ====================================================================================
112
+ # slice: bool=True
113
+ if args.options.slice.value is True:
114
+ args.slice = True
115
+ # elif args
116
+
117
+ # showfliers: bool=True
118
+ args.showfliers = True if args.options.showfliers.value or args.showfliers else False
119
+ # showmeans: bool=True
120
+ args.showmeans = True if args.options.showmeans.value or args.showmeans else False
121
+ # stest: bool=True
122
+ args.stest = True if args.options.statisticalTest.value or args.stest else False
123
+
124
+ # dpi: int
125
+ args.dpi = args.options.dpi.value if not args.options.dpi.is_default() else args.dpi
126
+
127
+ # out_dir: str=None
128
+ if not args.out_dir and args.options.outDir.is_set():
129
+ args.out_dir = args.options.outDir.value
130
+ if not os.path.exists(args.out_dir):
131
+ raise CE.OptionError(f"Provided {bold}outDir/out_dir{reset} is not exist.")
132
+
133
+ # file_name: str=None
134
+ if not args.file_name and args.options.fileName.is_set():
135
+ args.file_name = args.options.fileName.value
136
+
137
+ # title: str=None
138
+ if not args.title and args.options.title.is_set():
139
+ args.title = args.options.title.value
140
+
141
+ # ====================================================================================
142
+ # check num and source
143
+ if not args.source and args.options.source.value:
144
+ args.source = args.options.source.value
145
+
146
+ if not args.source and not args.options.source.value:
147
+ print(f"# WARNING: no data source provided, {bold}testing{reset} data is "
148
+ f"selected as default when plotting for {bold}experiments{reset}.")
149
+ args.source = 'testing'
150
+
151
+ if args.source in ('test', 'testing') or args.options.source.value in ('test', 'testing'):
152
+ if args.num != 'final' or args.options.numConfigurations.value != 'final':
153
+ print(f"# WARNING: when {bold}testing{reset} results selected as source, only {bold}final{reset} can be selected for {underline}num / numConfigurations{reset}")
154
+ args.num = 'final'
155
+
156
+
157
+ # check num and select
158
+ select = args.select
159
+ # select is provided only in options
160
+ if not args.select and args.options.selConfigurations.value:
161
+ select = args.options.selConfigurations.value
162
+ # select is not provided
163
+ elif not args.select and not args.options.selConfigurations.value:
164
+ if not args.num and not args.options.numConfigurations.value:
165
+ if name != 'heat':
166
+ raise CE.OptionError(f"No {underline}num/numConfigurations{reset} or "
167
+ f"{underline}select/selConfigurations{reset} provided for plotting {bold}experiments{reset}.")
168
+ if args.num == 'final' or args.options.numConfigurations.value == 'final':
169
+ select = args.data.elites
170
+ elif args.num == 'elites' or args.options.numConfigurations.value == 'elites':
171
+ if 'elites' in args.data.slice.columns:
172
+ tmp = [x[0] for x in args.data.slice.elites]
173
+ elif args.data.all_elites is not None:
174
+ tmp = args.data.all_elites
175
+ tmp.append(args.data.best_id)
176
+ select = list(dict.fromkeys(tmp))
177
+ elif args.num == 'all' or args.options.numConfigurations.value == 'all':
178
+ select = args.data.all_elites
179
+ elif args.num == 'else' or args.options.numConfigurations.value == 'else':
180
+ if not args.select or not args.options.selConfigurations.value:
181
+ raise CE.OptionError(f"{underline}select{reset} / {underline}selConfigurations{reset} "
182
+ f"must be provided when {bold}else{reset} is selected.")
183
+ args.select = select
184
+
185
+
186
+ if name == 'scat' and len(args.select) > 2:
187
+ del args.select[2:]
188
+ print(f"#\n# More than two configurations are selected for drawing scatter plot.\n"
189
+ f"# selected configurations: {args.select}")
190
+
191
+ @enforce_types
192
+ def _load_data(name: str, args: SimpleNamespace):
193
+ """load data"""
194
+ # load training results
195
+ if args.source == 'training':
196
+ src_quality = args.data.training.data.dropna()
197
+ # load test results
198
+ elif args.source in ('test', 'testing'):
199
+ src_quality = args.data.testing.data.dropna()
200
+ else:
201
+ raise CE.OptionError(f"No {underline}source{reset} data provided plotting for {bold}experiments{reset}.")
202
+
203
+ # select columns
204
+ sel_cols = ['experiment_id', 'configuration_id', 'instance_id', 'quality', 'end_time']
205
+ sel_pd = src_quality[sel_cols]
206
+
207
+ if name == 'heat':
208
+ return sel_pd
209
+ elif name == 'scat':
210
+ data = add_best(sel_pd)
211
+ # select configurations
212
+ # boxplot: may receive 2-dimensional select
213
+ data = data[data['configuration_id'].isin(np.unique(args.select+[0]))]
214
+ return data
215
+
216
+ # select configurations
217
+ new_pd = sel_pd[sel_pd['configuration_id'].isin(np.unique(args.select))]
218
+
219
+ # include slice information
220
+ new_pd = add_slice(data=new_pd, slice=args.data.slice)
221
+ new_pd = add_elites(data=new_pd, slice=args.data.slice, final=args.data.elites, select=args.select)
222
+
223
+ # remove extra information
224
+ data = new_pd.drop(columns=['end_time'])
225
+
226
+ # print(f"#\n# The original crace results:\n{new_pd}\n")
227
+
228
+ return data
229
+
230
+ @enforce_types
231
+ def _resolve_input(name: str, args: SimpleNamespace):
232
+ """
233
+ resolve input
234
+ """
235
+ new_data= None
236
+
237
+ _check_options(name=name, args=args)
238
+
239
+ new_data = _load_data(name=name, args=args)
240
+
241
+ return new_data
242
+
243
+ def plot_experiments(method: str, **kwargs):
244
+ """
245
+ entrance to call plotting function from file _draw.py
246
+ including try-except
247
+ """
248
+ # drawMethod: (boxplot, violinplot, parallelcoord, parallelcat, sunburst, pairplot, histplot, jointplot, heatmap)
249
+
250
+ func = dispatch_table_exps.get(method)
251
+
252
+ try:
253
+ if func: return func(**kwargs)
254
+ else: raise ValueError(f"Method '{method}' is not defined in dispatch table.")
255
+
256
+ except Exception as e:
257
+ print("\n! There was an error while plotting for parameters:")
258
+ if any(isinstance(e, cls) for cls in [x[1] for x in inspect.getmembers(CE, inspect.isclass)]):
259
+ print(f"! {e}")
260
+ else:
261
+ err = traceback.format_exc()
262
+ print(err)
263
+ if kwargs["_console"]: return None
264
+ sys.exit(1)
265
+
266
+ @export
267
+ @enforce_types
268
+ def qual_boxplot(
269
+ # required parameters
270
+ data: CraceResults=None,
271
+ options: CplotOptions=None,
272
+ select: list=None,
273
+ # slice: bool=None,
274
+ slice: Union[Literal['budget', 'b', 'B', 'e', 'E', 'experiment', 'time', 't', 'T'
275
+ ], bool, None]=None,
276
+ out_dir: str=None,
277
+ file_name: str=None,
278
+ title: str=None,
279
+ num: Literal['all', 'final', 'elites', None] = None,
280
+ source: Literal['training', 'test', 'testing', None] = None,
281
+ _console: bool=True,
282
+ # optional parameters
283
+ showfliers: bool=False,
284
+ showmeans: bool=False,
285
+ stest: bool=False,
286
+ dpi: int=800,
287
+ # specific parameters
288
+ x: str='configuration_id',
289
+ y: str='quality',
290
+ hue: str='configuration_id',
291
+ palette: str='vlag',
292
+ fontsize: int=12,
293
+ fliersize: float=.5,
294
+ show_ori: bool=False,
295
+ errorbar: Literal['ci', 'rpd'] = 'rpd',
296
+ check_elites: bool=False,
297
+ ):
298
+ """
299
+ Entrance to generate boxplots for quality comparison.
300
+
301
+ :param data: Object 'CraceResults' containing the quality results to be plotted.
302
+ :param options: Object 'CplotOptions' containing the plotting options.
303
+ :param select: Optional. A list of experiment, scenario, or method identifiers selected for plotting.
304
+ :param slice: Optional. Selects how the results are sliced before plotting. Supported values are 'budget' ('b', 'B'), 'experiment' ('e', 'E'), and 'time' ('t', 'T'). A boolean can also be used to enable or disable slicing.
305
+ :param out_dir: Output directory for the generated figure.
306
+ :param file_name: File name of the generated figure.
307
+ :param title: Optional title of the figure.
308
+ :param num: Selects the configurations to include in the plot: 'all', 'final', or 'elites'.
309
+ :param source: Selects the source of the results, either 'training' or 'test'/'testing'.
310
+ :param showfliers: Whether to display outliers in the boxplots.
311
+ :param showmeans: Whether to display the mean value.
312
+ :param stest: Whether to perform statistical tests between the compared groups.
313
+ :param dpi: Resolution of the output figure in dots per inch.
314
+ :param palette: Name of the color palette used for the boxplots.
315
+ :param fontsize: Font size used in the figure.
316
+ :param fliersize: Marker size used for outliers.
317
+ :param show_ori: Whether to show the original/default configuration.
318
+ :param errorbar: Type of error bar to display. Supported values are 'ci' for confidence interval and 'rpd' for relative percentage deviation.
319
+ :param check_elites: Whether to check and use elite configurations when generating the plot.
320
+ """
321
+ kwargs = locals()
322
+ plot_experiments(method='box', **kwargs)
323
+
324
+ @enforce_types
325
+ @register_exps(['boxplot', 'box'])
326
+ def _boxplot(
327
+ # required parameters
328
+ data: CraceResults=None,
329
+ options: CplotOptions=None,
330
+ select: list=None,
331
+ # slice: bool=None,
332
+ slice: Union[Literal['budget', 'b', 'B', 'e', 'E', 'experiment', 'time', 't', 'T'
333
+ ], bool, None]=None, out_dir: str=None,
334
+ file_name: str=None,
335
+ title: str=None,
336
+ num: Literal['all', 'final', 'elites', None] = None,
337
+ source: Literal['training', 'test', 'testing', None] = None,
338
+ _console: bool=True,
339
+ # optional parameters
340
+ showfliers: bool=False,
341
+ showmeans: bool=False,
342
+ stest: bool=False,
343
+ dpi: int=800,
344
+ # specific parameters
345
+ x: str='configuration_id',
346
+ y: str='quality',
347
+ hue: str='configuration_id',
348
+ palette: str='vlag',
349
+ fontsize: int=12,
350
+ fliersize: float=.5,
351
+ show_ori: bool=False,
352
+ errorbar: Literal['ci', 'rpd'] = 'rpd',
353
+ check_elites: bool=False,
354
+ ):
355
+ """drawing box plot"""
356
+
357
+ try:
358
+ import scikit_posthocs as sp
359
+ from statannotations.Annotator import Annotator
360
+ except ImportError as e:
361
+ raise e
362
+
363
+ init_plot_style(size=1.2)
364
+
365
+ args = SimpleNamespace(
366
+ data=safe_copy(data), options=safe_copy(options), num=num, select=select,
367
+ slice=slice, out_dir=out_dir, file_name=file_name, title=title,
368
+ source=source, console=_console,
369
+ # optional parameters
370
+ showfliers=showfliers, showmeans=showmeans, stest=stest,
371
+ dpi=dpi,
372
+ )
373
+
374
+ new_data = _resolve_input(name='box', args=args)
375
+
376
+ # if args.console: print_args(args)
377
+ # print(f"#\n# The original crace results:\n{new_data}\n")
378
+
379
+ # ====================================================================================
380
+ # draw plots for selected elite configurations
381
+ # in order to show the sampling
382
+ if not args.slice:
383
+ new_data = new_data[['experiment_id', 'configuration_id', 'instance_id', 'quality']].drop_duplicates()
384
+
385
+ results1 = new_data.groupby(['configuration_id']).mean().T.to_dict()
386
+ results2 = new_data.groupby(['configuration_id']).count().T.to_dict()
387
+ results_mean = {}
388
+ results_count = {}
389
+
390
+ for id, item in results2.items():
391
+ results_count[int(id)] = None
392
+ results_count[int(id)] = item['instance_id']
393
+
394
+ # mean based on Nins, when there is a tie on quality
395
+ for id, item in results1.items():
396
+ results_mean[int(id)] = None
397
+ results_mean[int(id)] = item['quality']
398
+
399
+ # Sort elite_ids based on results_count values
400
+ # 1. num of instances, increase
401
+ # 2. mean quality, decrease
402
+ results_count_sorted = {k:v for k, v in sorted(results_count.items(), key=lambda x: (x[1], -results_mean[x[0]]), reverse=False)}
403
+
404
+ elite_ids_sorted = [str(x) for x in results_count_sorted.keys()]
405
+ elite_ids_order = [int(x) for x in results_count_sorted.keys()]
406
+
407
+ min_quality = float("inf")
408
+ best_found = 0
409
+ pairs = [ k for k,v in results_count_sorted.items() if v==max(results_count_sorted.values())]
410
+ for id in pairs:
411
+ if results_mean[int(id)] <= min_quality:
412
+ min_quality = results_mean[int(id)]
413
+ best_found = id
414
+
415
+ best_final = data.best_id
416
+
417
+ # ====================================================================================
418
+ # update x-labels for 'training' results
419
+ x_labels = elite_ids_sorted.copy()
420
+ if args.source not in ('test', 'testing') and check_elites:
421
+ key_name = "-%(num)s-" % {"num": int(best_found)}
422
+ for i, x in enumerate(elite_ids_sorted):
423
+ if int(x) == int(best_final):
424
+ x = "*%(num)s*" % {"num": x}
425
+ if int(re.search(r'\d+', str(x)).group()) == int(best_found):
426
+ x = "-%(num)s-" % {"num": x}
427
+ elif int(x) == int(best_found):
428
+ x = key_name
429
+ x_labels[i] = x
430
+
431
+ print(f"# Best configurations: (num_instances, mean)")
432
+ print("# (number with *: final elite configuration from crace)")
433
+ print("# (number with -: configuration having the minimal average value on the most instances)")
434
+ for x in x_labels:
435
+ int_x = int(re.search(r'\d+', x).group())
436
+ print(f"#{x:>12}: ({results_count_sorted[int_x]:>2}, {results_mean[int_x]})", end="\n")
437
+
438
+ # ====================================================================================
439
+ # draw the plot
440
+ if num == 'all':
441
+ fig, ax = plt.subplots(figsize=(16,12))
442
+ else:
443
+ fig, ax = plt.subplots() # pylint: disable=undefined-variable
444
+ sns.boxplot(x='configuration_id', y='quality', data=new_data,
445
+ whis=0.5, showfliers=args.showfliers, fliersize=fliersize,
446
+ showmeans=args.showmeans,
447
+ meanprops={"marker": "x",
448
+ "markeredgecolor": "black",
449
+ "markersize": 1,
450
+ "linewidth": 2*fliersize},
451
+ width=0.4, linewidth=2*fliersize,
452
+ palette=palette,
453
+ boxprops=dict(facecolor="white"),
454
+ medianprops={"linewidth": 2*fliersize, "color": "black"},
455
+ whiskerprops={"linestyle": "--", "linewidth": 2*fliersize},
456
+ # flierprops={"marker": "o", "markersize": fliersize, 'markerfacecolor': 'black', 'markeredgecolor': 'black'},
457
+ hue='configuration_id', legend=False,
458
+ order=elite_ids_order,
459
+ # notch=True,
460
+ ax=ax,
461
+ )
462
+
463
+ if errorbar == 'ci':
464
+ _legend_ci(data=new_data, fliersize=fliersize, fontsize=fontsize, ax=ax, fig=fig, order=elite_ids_order)
465
+ elif errorbar == 'rpd':
466
+ _legend_rpd(data=new_data, fliersize=fliersize, fontsize=fontsize, ax=ax, fig=fig, order=elite_ids_order)
467
+
468
+ if show_ori:
469
+ sns.stripplot(x='configuration_id', y='quality', data=new_data,
470
+ color='green', size=4*fliersize, jitter=False,
471
+ order=elite_ids_order, ax=ax)
472
+
473
+ if args.stest:
474
+ _do_stest(data=new_data, args=args, pairs=pairs)
475
+
476
+ # add p-value for the configurations who have the most instances
477
+ if len(pairs) > 1:
478
+ if (int(best_final) != best_found and
479
+ int(best_final) in pairs):
480
+ pair = (int(best_found), int(best_final))
481
+ else:
482
+ pair = (int(elite_ids_sorted[-2]), int(elite_ids_sorted[-1]))
483
+ pairs_results = new_data.loc[new_data['configuration_id'].isin(pair)].copy()
484
+
485
+ # p1 = sp.posthoc_wilcoxon(pairs_results, val_col='quality', group_col='configuration_id')
486
+ # p_values after multiple test correction
487
+ p2 = sp.posthoc_wilcoxon(pairs_results, val_col='quality', group_col='configuration_id',
488
+ p_adjust='fdr_bh')
489
+ p2_4 = p2.round(4)
490
+
491
+ print("#\n# Adjusted p-values of the last two elite configurations:")
492
+ print(p2_4)
493
+
494
+ annotator = Annotator(ax, pairs=[pair], data=new_data, order=elite_ids_order, x='configuration_id', y='quality')
495
+ annotator.configure(test='Wilcoxon', text_format='simple', comparisons_correction='fdr_bh',
496
+ show_test_name=False, line_width=1)
497
+ annotator.apply_and_annotate()
498
+
499
+ # update labels / ticks
500
+ ticks = ax.get_xticks()
501
+ xticklabels = [t.get_text() for t in ax.get_xticklabels()]
502
+ ax.set_xticks(ticks)
503
+ if len(x_labels) > 12:
504
+ ax.set_xticklabels(x_labels, rotation=90)
505
+ else:
506
+ ax.set_xticklabels(x_labels, rotation=0)
507
+ if not args.title:
508
+ ax.set_xlabel(None)
509
+ elif not args.console:
510
+ ax.set_xlabel(args.title)
511
+ else:
512
+ x.set_xlabel()
513
+
514
+ if not data.options.capping.value:
515
+ ax.set_ylabel('quality')
516
+ else:
517
+ ax.set_ylabel('runtime')
518
+ plt.xticks()
519
+ plt.yticks()
520
+
521
+ _legend_ins(ax=ax, ticks=ticks, labels=results_count_sorted)
522
+
523
+
524
+ else:
525
+ # test
526
+ num = len(new_data.slice_finished.unique())
527
+ if num == 1:
528
+ sns.boxplot(x=x, y=y, data=new_data,
529
+ hue=hue, legend=None,
530
+ whis=0.5, showfliers=args.showfliers, fliersize=fliersize,
531
+ showmeans=args.showmeans,
532
+ meanprops={"marker": "x",
533
+ "markeredgecolor": "black",
534
+ "markersize": 1,
535
+ "linewidth": 2*fliersize},
536
+ width=0.4, linewidth=2*fliersize,
537
+ palette=palette,
538
+ boxprops=dict(facecolor="white"),
539
+ medianprops={"linewidth": 2*fliersize, "color": "black"},
540
+ whiskerprops={"linestyle": "--", "linewidth": 2*fliersize},
541
+ ax=ax,
542
+ )
543
+
544
+ ticks = ax.get_xticks()
545
+ fig = ax.get_figure()
546
+ xticklabels = [int(t.get_text()) for t in ax.get_xticklabels()]
547
+
548
+ if errorbar == 'ci':
549
+ _legend_ci(data=new_data, fliersize=fliersize, fontsize=fontsize, ax=ax, fig=fig, order=xticklabels, x=x, y=y, hue=hue)
550
+ elif errorbar == 'rpd':
551
+ _legend_rpd(data=new_data, fliersize=fliersize, fontsize=fontsize, ax=ax, fig=fig, order=xticklabels, x=x, y=y, hue=hue)
552
+
553
+ if show_ori:
554
+ sns.stripplot(x=x, y=y, data=new_data,
555
+ color='green', size=4*fliersize, jitter=False,
556
+ order=xticklabels, ax=ax)
557
+
558
+ if not data.options.capping.value:
559
+ ax.set_ylabel('quality')
560
+ else:
561
+ ax.set_ylabel('runtime')
562
+ plt.xticks()
563
+ plt.yticks()
564
+
565
+ # training
566
+ else:
567
+ col_name = 'slice_elites' if args.select is None else 'group'
568
+ col_wrap = _auto_col_wrap(num)
569
+
570
+ p = sns.FacetGrid(data=new_data, col=col_name, col_wrap=num, legend_out=False,
571
+ sharex=False, sharey=True, height=3, aspect=0.7)
572
+ p.map_dataframe(func=sns.boxplot, x=x, y=y,
573
+ hue=hue, legend=True,
574
+ whis=0.5, showfliers=args.showfliers, fliersize=fliersize,
575
+ showmeans=args.showmeans,
576
+ meanprops={"marker": "x",
577
+ "markeredgecolor": "black",
578
+ "markersize": 1,
579
+ "linewidth": 2*fliersize},
580
+ width=0.4, linewidth=2*fliersize,
581
+ palette=palette,
582
+ boxprops=dict(facecolor="white"),
583
+ medianprops={"linewidth": 2*fliersize, "color": "black"},
584
+ whiskerprops={"linestyle": "--", "linewidth": 2*fliersize},
585
+ )
586
+
587
+ left_axes = []
588
+ top_axes = []
589
+ idx = 1
590
+ for ax, (_, subdata) in zip(p.axes.flatten(), p.facet_data()):
591
+ print(f"#\n# The original crace results for slice {idx}:\n{subdata}\n")
592
+
593
+ if ax.get_subplotspec().colspan.start == 0:
594
+ left_axes.append(ax)
595
+
596
+ if ax.get_subplotspec().rowspan.start == 0:
597
+ top_axes.append(ax)
598
+
599
+ ticks = ax.get_xticks()
600
+ xticklabels = [int(t.get_text()) for t in ax.get_xticklabels()]
601
+
602
+ if errorbar == "ci":
603
+ _legend_ci(data=subdata, fliersize=fliersize, fontsize=fontsize, ax=ax, fig=ax.figure, add_legend=False, order=xticklabels, x=x, y=y, hue=hue)
604
+
605
+ elif errorbar == "rpd":
606
+ _legend_rpd(data=subdata, fliersize=fliersize, fontsize=fontsize, ax=ax, fig=ax.figure, add_legend=False, order=xticklabels, x=x, y=y, hue=hue)
607
+
608
+ if show_ori:
609
+ sns.stripplot(x=x, y=y, data=subdata,
610
+ color='green', size=4*fliersize, jitter=False,
611
+ order=xticklabels, ax=ax)
612
+
613
+ ax.set_xlabel("")
614
+ ax.set_ylabel("")
615
+ idx += 1
616
+
617
+ fig = p.figure
618
+
619
+ fig.canvas.draw()
620
+ renderer = fig.canvas.get_renderer()
621
+
622
+ # Plot area
623
+ axes_bboxes = [
624
+ ax.get_position()
625
+ for ax in p.axes.flatten()
626
+ if ax.get_visible()
627
+ ]
628
+
629
+ x_min = min(b.x0 for b in axes_bboxes)
630
+ x_max = max(b.x1 for b in axes_bboxes)
631
+ y_min = min(b.y0 for b in axes_bboxes)
632
+
633
+
634
+ # X label
635
+ supx = fig.supxlabel("configuration_id")
636
+ _, y = supx.get_position()
637
+ supx.set_position(((x_min + x_max) / 2, y))
638
+
639
+
640
+ # Y tick labels
641
+ tick_bboxes = [
642
+ label.get_window_extent(renderer=renderer)
643
+ for ax in left_axes
644
+ for label in ax.get_yticklabels()
645
+ if label.get_visible() and label.get_text()
646
+ ]
647
+
648
+ if tick_bboxes:
649
+ tick_bbox = Bbox.union(tick_bboxes)
650
+ tick_bbox = tick_bbox.transformed(fig.transFigure.inverted())
651
+ x = tick_bbox.x0 - 0.025
652
+ else:
653
+ x = x_min - 0.05
654
+
655
+ y_min = min(ax.get_position().y0 for ax in left_axes)
656
+ y_max = max(ax.get_position().y1 for ax in left_axes)
657
+
658
+ fig.text(
659
+ x,
660
+ (y_min + y_max) / 2,
661
+ 'quality' if not data.options.capping.value else 'runtime',
662
+ rotation=90,
663
+ ha="center",
664
+ va="center",
665
+ )
666
+
667
+ # add legend
668
+ text = 'median (95% CI)' if errorbar == 'ci' else 'median ± RPD'
669
+ _add_legend(text=text, fontsize=fontsize, glob=True, fig=fig, top=top_axes, left=left_axes)
670
+
671
+ if args.stest:
672
+ plog = args.file_name
673
+
674
+ # p_values
675
+ p1 = sp.posthoc_wilcoxon(data, val_col='quality', group_col='exp_name')
676
+ # p_values after multiple test correction
677
+ p2 = sp.posthoc_wilcoxon(data, val_col='quality', group_col='exp_name',
678
+ p_adjust='fdr_bh')
679
+ print("Original p_values caculated by 'Wilcoxon':\n", p1)
680
+ print("New p_values corrected by 'fdr_bh':\n", p2)
681
+
682
+ with open(args.out_dir + "/" + plog + '.log', 'w') as f1:
683
+ print("Original p_values caculated by 'Wilcoxon':\n", p1, file=f1)
684
+ print("\nNew p_values corrected by 'fdr_bh':\n", p2, file=f1)
685
+ print("\n", file=f1)
686
+
687
+ order = []
688
+ pairs = []
689
+ p_values = []
690
+ for x in data['exp_name'].unique():
691
+ order.append(x)
692
+ i = 0
693
+ for x in order[:-1]:
694
+ i += 1
695
+ for y in order[i:]:
696
+ pairs.append((x,y))
697
+ p_values.append(p2.loc[x, y])
698
+
699
+ annotator = Annotator(ax, pairs=pairs, order=elite_ids_order,
700
+ data=data, x='exp_name', y='quality')
701
+ annotator.configure(test='Wilcoxon', text_format='star', comparisons_correction='fdr_bh',
702
+ line_width=0.5)
703
+
704
+ with open(args.out_dir + "/" + plog + '.log', 'a') as f1:
705
+ original_stdout = sys.stdout
706
+ sys.stdout = f1
707
+
708
+ try:
709
+ annotator.apply_and_annotate()
710
+ finally:
711
+ sys.stdout = original_stdout
712
+
713
+ # plot = fig.get_figure()
714
+ # if not _console:
715
+ # plot.savefig(f"{args.out_dir}/{args.file_name}.png", dpi=args.dpi)
716
+ # print("# {} has been saved in {}.".format(args.file_name, args.out_dir))
717
+ plt.show()
718
+
719
+
720
+ @export
721
+ @enforce_types
722
+ def qual_scatter(
723
+ # required parameters
724
+ data: CraceResults=None,
725
+ options: CplotOptions=None,
726
+ select: list=None,
727
+ # slice: bool=None,
728
+ slice: Union[Literal['budget', 'b', 'B', 'e', 'E', 'experiment', 'time', 't', 'T'
729
+ ], bool, None]=None,
730
+ out_dir: str=None,
731
+ file_name: str=None,
732
+ title: str=None,
733
+ num: Literal['all', 'final', 'elites', None] = None,
734
+ source: Literal['training', 'test', 'testing', None] = None,
735
+ _console: bool=True,
736
+ # optional parameters
737
+ showfliers: bool=False,
738
+ showmeans: bool=False,
739
+ stest: bool=False,
740
+ dpi: int=800,
741
+ # specific parameters
742
+ xid: int=None,
743
+ yid: int=None,
744
+ fontsize: int=12,
745
+ fliersize: float=.5,
746
+ fillna: bool=False,
747
+ penalty: float=1.2,
748
+ rpd: bool=True,
749
+ instance_ids: list=None
750
+ ):
751
+ """
752
+ Entrance to generate scatter plots for quality comparison.
753
+
754
+ :param data: Object CraceResults that must be provided.
755
+ :param options: Object CplotOptions that must be provided.
756
+ :param select: Optional. A list of experiment, scenario, or method identifiers selected for plotting.
757
+ :param slice: Optional. Specifies how the quality data are sliced before plotting. Supported values are 'budget'/'b'/'B', 'experiment'/'e'/'E', 'time'/'t'/'T', or a boolean value.
758
+ :param out_dir: Output directory for the generated figure.
759
+ :param file_name: File name of the generated figure.
760
+ :param title: Optional title of the figure.
761
+ :param num: Selects the configurations to include in the plot: 'all', 'final', or 'elites'.
762
+ :param source: Selects the source of the results, either 'training' or 'test'/'testing'.
763
+ :param showfliers: Whether to display outliers in the scatter plot.
764
+ :param showmeans: Whether to display the mean value.
765
+ :param stest: Whether to perform statistical tests between the compared groups.
766
+ :param dpi: Resolution of the output figure in dots per inch.
767
+ :param xid: Identifier of the quality measure plotted on the x-axis.
768
+ :param yid: Identifier of the quality measure plotted on the y-axis.
769
+ :param fontsize: Font size used in the figure.
770
+ :param fliersize: Marker size used for outliers.
771
+ :param fillna: Whether to fill missing values before plotting.
772
+ :param penalty: Penalty factor applied to missing or invalid quality values.
773
+ :param rpd: Whether to use relative percentage deviation for the plotted quality values.
774
+ :param instance_ids: Optional list of instance identifiers to include in the plot.
775
+ """
776
+ kwargs = locals()
777
+ plot_experiments(method='scat', **kwargs)
778
+
779
+ @enforce_types
780
+ @register_exps(['scatter', 'scat'])
781
+ def _scatter(
782
+ # required parameters
783
+ data: CraceResults=None,
784
+ options: CplotOptions=None,
785
+ select: list=None,
786
+ # slice: bool=None,
787
+ slice: Union[Literal['budget', 'b', 'B', 'e', 'E', 'experiment', 'time', 't', 'T'
788
+ ], bool, None]=None,
789
+ out_dir: str=None,
790
+ file_name: str=None,
791
+ title: str=None,
792
+ num: Literal['all', 'final', 'elites', None] = None,
793
+ source: Literal['training', 'test', 'testing', None] = None,
794
+ _console: bool=True,
795
+ # optional parameters
796
+ showfliers: bool=False,
797
+ showmeans: bool=False,
798
+ stest: bool=False,
799
+ dpi: int=800,
800
+ # specific parameters
801
+ xid: int=None,
802
+ yid: int=None,
803
+ fontsize: int=12,
804
+ fliersize: float=.5,
805
+ fillna: bool=False,
806
+ penalty: float=1.2,
807
+ rpd: bool=True,
808
+ instance_ids: list=None
809
+ ):
810
+ """drawing scatter plot"""
811
+ if (select is None and options.selConfigurations.value is None) and (
812
+ xid is not None and yid is not None
813
+ ):
814
+ select = [xid, yid]
815
+
816
+ args = SimpleNamespace(
817
+ data=safe_copy(data), options=safe_copy(options), num=num, select=select,
818
+ slice=slice, out_dir=out_dir, file_name=file_name, title=title,
819
+ source=source, console=_console,
820
+ # optional parameters
821
+ showfliers=showfliers, showmeans=showmeans, stest=stest, dpi=dpi,
822
+ )
823
+
824
+ new_data = _resolve_input(name='scat', args=args)
825
+
826
+ # if args.console: print_args(args)
827
+ # print(f"#\n# The original crace results:\n{new_data}\n")
828
+
829
+ sel_ins = new_data.loc[
830
+ new_data["configuration_id"].isin(args.select),
831
+ "instance_id"].unique()
832
+
833
+ sel_data = new_data[~(
834
+ (new_data["configuration_id"] == 0) &
835
+ (~new_data["instance_id"].isin(sel_ins)))]
836
+
837
+ pivot = sel_data.pivot(
838
+ index="instance_id",
839
+ columns="configuration_id",
840
+ values="quality")
841
+
842
+ if not fillna:
843
+ pivot_data = pivot.dropna()
844
+ else:
845
+ fill_num = pivot.max().max() * penalty
846
+ pivot_data = pivot.fillna(fill_num)
847
+ print(f"# WARNING: fill NaN with {bold}{fill_num}{reset}, calculated by: {bold}data.max().max() * penalty({penalty}){reset}\n")
848
+
849
+ pivot_show = pivot_data.rename(columns={0: "bests"})
850
+ mean_row = pivot_show.mean().to_frame().T
851
+ mean_row.index = ["mean"]
852
+ pivot_show = pd.concat([pivot_show, mean_row])
853
+ print(pivot_show)
854
+ print(f"Note: {bold}bests{reset} represents an oracle configuration, constructed by selecting, "
855
+ f"for each instance,\n\t the configuration that achieves the best performance.\n")
856
+
857
+ if xid is None and yid is None:
858
+ xid, yid = args.select
859
+
860
+ orig_instance_ids = pivot_data.index.astype(int)
861
+
862
+ if rpd:
863
+ pivot_data = _scatter_rpd(pivot_data)
864
+ xlab = f"RPD (%) of configuration {xid}"
865
+ ylab = f"RPD (%) of configuration {yid}"
866
+ else:
867
+ xlab = f"Cost of configuration {xid}"
868
+ ylab = f"Cost of configuration {yid}"
869
+
870
+ x_data = pivot_data[xid]
871
+ y_data = pivot_data[yid]
872
+
873
+ mask = x_data.notna() & y_data.notna()
874
+ if not mask.any():
875
+ raise ValueError("No instance has data for both configurations")
876
+
877
+ x_data = x_data[mask]
878
+ y_data = y_data[mask]
879
+ instances = orig_instance_ids[mask]
880
+
881
+ # find better
882
+ best = np.full(len(x_data), "equal", dtype=object)
883
+ best[x_data < y_data] = "conf1"
884
+ best[x_data > y_data] = "conf2"
885
+
886
+ # ---- instance names ----
887
+ if instance_ids is None:
888
+ instance_ids = instances
889
+ elif callable(instance_ids):
890
+ instance_ids = [instance_ids(x) for x in instances]
891
+ else:
892
+ if len(instance_ids) != len(pivot_data):
893
+ raise ValueError("`instance_ids` must have same length as experiments")
894
+ instance_ids = np.asarray(instance_ids)[mask]
895
+
896
+ # new dataframe
897
+ df = pd.DataFrame({
898
+ "conf1": x_data.values,
899
+ "conf2": y_data.values,
900
+ "instance": instance_ids,
901
+ "best": best,
902
+ })
903
+
904
+ _draw_scatter(df, xlab, ylab)
905
+
906
+ def _scatter_rpd(df, ref_id=0):
907
+ """calculate rpd information for selected configurations on instances"""
908
+ if ref_id in df.columns:
909
+ best = df[ref_id]
910
+ else:
911
+ best = df.min(axis=1)
912
+ return (df.sub(best, axis=0)).div(best, axis=0)
913
+
914
+ def _draw_scatter(df, xlab, ylab):
915
+ init_plot_style()
916
+ fig, ax = plt.subplots()
917
+
918
+ # colors = {
919
+ # "conf1": "#0055CC",
920
+ # "conf2": "#C41700",
921
+ # "equal": "darkgray",
922
+ # }
923
+ # for key, g in df.groupby("best"):
924
+ # ax.scatter(
925
+ # g["conf1"],
926
+ # g["conf2"],
927
+ # s=40,
928
+ # color=colors[key],
929
+ # label=key,
930
+ # alpha=0.9,
931
+ # )
932
+
933
+ import matplotlib.colors as mcolors
934
+
935
+ norm = mcolors.Normalize(
936
+ vmin=df["instance"].min(),
937
+ vmax=df["instance"].max()
938
+ )
939
+
940
+ for key, g in df.groupby("best"):
941
+ ax.scatter(
942
+ g["conf1"],
943
+ g["conf2"],
944
+ c=g["instance"], # 用 instance 控制颜色
945
+ cmap="viridis",
946
+ norm=norm,
947
+ s=40,
948
+ alpha=0.85,
949
+ label=key,
950
+ edgecolors="none"
951
+ )
952
+
953
+ plt.colorbar(
954
+ plt.cm.ScalarMappable(norm=norm, cmap="viridis"),
955
+ ax=ax,
956
+ label="Instance ID"
957
+ )
958
+
959
+ # y = x
960
+ lim = max(df["conf1"].max(), df["conf2"].max())
961
+ ax.plot([0, lim], [0, lim], color="lightgray", linewidth=1.5)
962
+
963
+ ax.set_xlabel(xlab)
964
+ ax.set_ylabel(ylab)
965
+
966
+ ax.legend().remove()
967
+
968
+ plt.tight_layout()
969
+ plt.show()
970
+
971
+
972
+ @export
973
+ @enforce_types
974
+ def qual_heatmap(
975
+ # required parameters
976
+ data: CraceResults=None,
977
+ options: CplotOptions=None,
978
+ select: list=None,
979
+ # slice: bool=None,
980
+ slice: Union[Literal['budget', 'b', 'B', 'e', 'E', 'experiment', 'time', 't', 'T'
981
+ ], bool, None]=None,
982
+ out_dir: str=None,
983
+ file_name: str=None,
984
+ title: str=None,
985
+ num: Literal['all', 'final', 'elites', None] = None,
986
+ source: Literal['training', 'test', 'testing', None] = None,
987
+ _console: bool=True,
988
+ # optional parameters
989
+ showfliers: bool=False,
990
+ showmeans: bool=False,
991
+ stest: bool=False,
992
+ dpi: int=800,
993
+ # specific parameters
994
+ fontsize: int=12,
995
+ fliersize: float=.5,
996
+ colorscale: str="Viridis",
997
+ return_fig: bool=False,
998
+ ):
999
+ """
1000
+ Entrance to call quality heatmap in python console
1001
+
1002
+ :param data: Object CraceResults that must be provided.
1003
+ :param options: Object CplotOptions that must be provided.
1004
+ :param select: A list of instance names or identifiers selected for plotting.
1005
+ :param slice: Optional. Specifies how the quality data are sliced before plotting. Supported values are 'budget'/'b'/'B', 'experiment'/'e'/'E', 'time'/'t'/'T', or a boolean value.
1006
+ :param out_dir: Output directory for saving the generated figure.
1007
+ :param file_name: File name of the generated figure.
1008
+ :param title: Title of the heatmap.
1009
+ :param num: Optional. Specifies the configurations included in the heatmap. Supported values are 'all', 'final', and 'elites'.
1010
+ :param source: Optional. Specifies the source of the quality data. Supported values are 'training' and 'test'/'testing'.
1011
+ :param showfliers: Boolean used to enable/disable showing outliers.
1012
+ :param showmeans: Boolean used to enable/disable showing mean values.
1013
+ :param stest: Boolean used to enable/disable statistical testing.
1014
+ :param dpi: Resolution of the generated figure in dots per inch.
1015
+ :param fontsize: Font size used in the heatmap.
1016
+ :param fliersize: Size of the outlier markers.
1017
+ :param colorscale: A string of palette name used for plotting.
1018
+ """
1019
+ kwargs = locals()
1020
+ if return_fig:
1021
+ return plot_experiments(method='heat', **kwargs)
1022
+ else:
1023
+ plot_experiments(method='heat', **kwargs)
1024
+
1025
+ @enforce_types
1026
+ @register_exps(['heatmap', 'heat'])
1027
+ def _heatmap(
1028
+ # required parameters
1029
+ data: CraceResults=None,
1030
+ options: CplotOptions=None,
1031
+ select: list=None,
1032
+ # slice: bool=None,
1033
+ slice: Union[Literal['budget', 'b', 'B', 'e', 'E', 'experiment', 'time', 't', 'T'
1034
+ ], bool, None]=None,
1035
+ out_dir: str=None,
1036
+ file_name: str=None,
1037
+ title: str=None,
1038
+ num: Literal['all', 'final', 'elites', None] = None,
1039
+ source: Literal['training', 'test', 'testing', None] = None,
1040
+ _console: bool=True,
1041
+ # optional parameters
1042
+ showfliers: bool=False,
1043
+ showmeans: bool=False,
1044
+ stest: bool=False,
1045
+ dpi: int=800,
1046
+ # specific parameters
1047
+ fontsize: int=12,
1048
+ fliersize: float=.5,
1049
+ colorscale: str="Viridis",
1050
+ return_fig: bool=False,
1051
+ ):
1052
+ try:
1053
+ import plotly.graph_objects as go
1054
+ except ImportError as e:
1055
+ raise e
1056
+
1057
+ """drawing scatter plot"""
1058
+ args = SimpleNamespace(
1059
+ data=safe_copy(data), options=safe_copy(options), num=num, select=select,
1060
+ slice=slice, out_dir=out_dir, file_name=file_name, title=title,
1061
+ source=source, console=_console,
1062
+ # optional parameters
1063
+ showfliers=showfliers, showmeans=showmeans, stest=stest, dpi=dpi,
1064
+ )
1065
+
1066
+ new_data = _resolve_input(name='heat', args=args)
1067
+
1068
+ # if args.console: print_args(args)
1069
+ # print(f"#\n# The original crace results:\n{new_data}\n")
1070
+
1071
+ pivot = new_data.pivot(
1072
+ index="instance_id",
1073
+ columns="configuration_id",
1074
+ values="quality"
1075
+ )
1076
+
1077
+ fig = go.Figure(
1078
+ data=go.Heatmap(
1079
+ z=pivot.values,
1080
+ x=pivot.columns,
1081
+ y=pivot.index,
1082
+ colorscale=colorscale,
1083
+ colorbar=dict(title="Quality"),
1084
+ zmin=pivot.min().min(),
1085
+ zmax=pivot.max().max(),
1086
+ )
1087
+ )
1088
+
1089
+ fig.update_layout(
1090
+ xaxis_title="Configuration IDs",
1091
+ yaxis_title="Instance IDs",
1092
+ template="plotly_white",
1093
+ )
1094
+
1095
+ if return_fig: return fig
1096
+ else: fig.show()
1097
+
1098
+ def _do_stest(data: pd.DataFrame, args, pairs):
1099
+ """add stest information for boxplot"""
1100
+
1101
+ try:
1102
+ from statsmodels.formula.api import ols
1103
+ import statsmodels.api as sm
1104
+ import scikit_posthocs as sp
1105
+ import scipy.stats as stats
1106
+ import logging
1107
+ except ImportError as e:
1108
+ raise e
1109
+
1110
+ l = logging.getLogger('st_log')
1111
+ filehandler = logging.FileHandler(args.out_dir + "/" + args.file_name + '.log', mode='w')
1112
+ filehandler.setLevel(0)
1113
+ streamhandler = logging.StreamHandler()
1114
+ l.setLevel(logging.DEBUG)
1115
+ l.addHandler(filehandler)
1116
+ l.addHandler(streamhandler)
1117
+
1118
+ data = data.loc[data['configuration_id'].isin(pairs)].copy()
1119
+
1120
+ ############################# CHECK RESULTS #############################
1121
+ # Shapiro-Wilk Test #
1122
+ # LEVENE #
1123
+ # ANOVA #
1124
+ # Kruskal-Wallis H Test #
1125
+ #########################################################################
1126
+
1127
+ # avg for each configuration
1128
+ data_groups = [data['quality'][data['configuration_id'] == conf] for conf in data['configuration_id'].unique()]
1129
+ print("data_groups: ", data_groups)
1130
+
1131
+ # Shapiro-Wilk Test
1132
+ # H0 hypothesis: normality (normal distribution)
1133
+ shapiro_string = ''
1134
+ stat_s = p_s = []
1135
+ for conf in data['configuration_id'].unique():
1136
+ data_group = data[data['configuration_id'] == conf]['quality']
1137
+ ss, ps = stats.shapiro(data_group)
1138
+ stat_s.append(ss)
1139
+ p_s.append(ps)
1140
+ shapiro_string += 'Shapiro-Wilk Test for configuration {}, Statistic: {:.4f}, p-value: {:.4f}\n'.format(conf, ss, ps)
1141
+ l.debug(f'\nShapiro-Wilk Test - H0 hypothesis: normality (0.05)\n{shapiro_string}')
1142
+
1143
+ # do levene
1144
+ # H0 hypothesis: homogeneity of variance (方差齐性)
1145
+ stat_l, p_l = stats.levene(*data_groups)
1146
+ l.debug('\nLevene’s Test - H0 hypothesis: homogeneity of variance (0.05)\n' \
1147
+ 'stat_l: {:.4f}, p-value: {:.4f}\n'.format(stat_l, p_l))
1148
+
1149
+ # check the results from Shapiro-Wilk Test and levene
1150
+ KW_test = ANOVA_test = False
1151
+ if p_l < 0.05 or any(x<0.05 for x in p_s):
1152
+ KW_test = True
1153
+ else:
1154
+ ANOVA_test = True
1155
+
1156
+ if ANOVA_test:
1157
+ # simulate ANOVA
1158
+ # H0 hypothesis: same mean values
1159
+ model = ols('quality ~ C(configuration_id)', data=data).fit()
1160
+ anova_results = sm.stats.anova_lm(model, typ=2) # Type 2 ANOVA DataFrame
1161
+ l.debug(f'\nANOVA_results - H0 hypothesis: all configurations have the same mean values\n{anova_results}')
1162
+
1163
+ if KW_test:
1164
+ # do Kruskal-Wallis H
1165
+ # H0 hypothesis: same medians
1166
+ stat_k, p_k = stats.kruskal(*data_groups)
1167
+ l.debug('Kruskal-Wallis Test - H0 hypothesis: all configurations have the same medians (0.05)\n' \
1168
+ 'Statistic: {:.4f}, p-value: {:.4f}'.format(stat_k, p_k))
1169
+
1170
+ ############################# POSTHOC TEST ##############################
1171
+ # posthoc_dunn #
1172
+ # posthoc_mannwhitney #
1173
+ #########################################################################
1174
+
1175
+ # # # Dunn:
1176
+ # p1 = sp.posthoc_dunn(data, val_col='quality', group_col='configuration_id')
1177
+ # # p_values after multiple test correction
1178
+ # p2 = sp.posthoc_dunn(data, val_col='quality', group_col='configuration_id',
1179
+ # p_adjust='fdr_bh')
1180
+ # p2_4 = p2.round(4)
1181
+
1182
+ # l.debug(f"\nOriginal p_values caculated by 'dunn':\n{p1}")
1183
+ # l.debug(f"\nNew p_values corrected by 'fdr_bh':\n{p2}")
1184
+ # l.debug(f"\nNew rounded p_values:\n{p2_4}")
1185
+
1186
+ # Wilcoxon rank-sum test
1187
+ p1 = sp.posthoc_mannwhitney(data, val_col='quality', group_col='configuration_id')
1188
+ # p_values after multiple test correction
1189
+ p2 = sp.posthoc_mannwhitney(data, val_col='quality', group_col='configuration_id',
1190
+ p_adjust='fdr_bh')
1191
+ p2_4 = p2.round(4)
1192
+
1193
+ l.debug(f"\nOriginal p_values caculated by 'mannwhitney (Wilcoxon rank-sum test)':\n{p1}")
1194
+ l.debug(f"\nNew p_values corrected by 'fdr_bh':\n{p2}")
1195
+ l.debug(f"\nNew rounded p_values:\n{p2_4}")
1196
+
1197
+ ############################# POSTHOC TEST ##############################
1198
+ # posthoc_wilcoxon #
1199
+ #########################################################################
1200
+
1201
+ # Wilcoxon signed-rank test
1202
+ p1 = sp.posthoc_wilcoxon(data, val_col='quality', group_col='configuration_id')
1203
+ # p_values after multiple test correction
1204
+ p2 = sp.posthoc_wilcoxon(data, val_col='quality', group_col='configuration_id',
1205
+ p_adjust='fdr_bh')
1206
+ p2_4 = p2.round(4)
1207
+ l.debug(f"\n############################# Wilcoxon Signed-rank Test ##############################")
1208
+ l.debug(f"\nOriginal p_values caculated by 'Wilcoxon':\n{p1}")
1209
+ l.debug(f"\nNew p_values corrected by 'fdr_bh':\n{p2}")
1210
+ l.debug(f"\nNew rounded p_values:\n{p2_4}")
1211
+
1212
+ def _auto_col_wrap(n):
1213
+ """
1214
+ Given number of subplots n,
1215
+ find the closest factor pair (a, b) with a*b = n,
1216
+ and return the larger one (b) as col_wrap.
1217
+ """
1218
+ best_pair = (1, n)
1219
+ min_diff = n - 1
1220
+
1221
+ for i in range(1, int(math.sqrt(n)) + 1):
1222
+ if n % i == 0:
1223
+ j = n // i
1224
+ if abs(j - i) < min_diff:
1225
+ best_pair = (i, j)
1226
+ min_diff = abs(j - i)
1227
+
1228
+ return max(best_pair)
1229
+
1230
+ def _add_legend(text, fontsize, ax=None, fig=None, glob=False, top:list=None, left:list=None):
1231
+ legend_element = mlines.Line2D([], [],
1232
+ color='red', marker='_', linestyle='None', markersize=.25*fontsize,
1233
+ label=text)
1234
+
1235
+ plt.tight_layout()
1236
+
1237
+ if not glob:
1238
+ offset_text = ax.yaxis.get_offset_text()
1239
+ if fig is None or not hasattr(fig, "canvas"):
1240
+ fig = ax.figure
1241
+ fig.canvas.draw()
1242
+ renderer = fig.canvas.get_renderer()
1243
+
1244
+ bbox = offset_text.get_window_extent(renderer=renderer)
1245
+
1246
+ fig_box = ax.get_position()
1247
+
1248
+ y_disp = (bbox.y0 + bbox.y1) / 2
1249
+ _, offset_y = fig.transFigure.inverted().transform((0, y_disp))
1250
+ # _, offset_y = ax.transAxes.inverted().transform((0, y_disp))
1251
+
1252
+ offset_y = max(0.02, min(1.0, offset_y))
1253
+
1254
+ offset_x = 0.99
1255
+
1256
+ else:
1257
+ supx = fig._supxlabel
1258
+ fig.canvas.draw()
1259
+ renderer = fig.canvas.get_renderer()
1260
+ bbox = supx.get_window_extent(renderer=renderer)
1261
+
1262
+ x_max = max(ax.get_position().x1 for ax in top)
1263
+ y_disp = (bbox.y0 + bbox.y1)/2
1264
+
1265
+ offset_x = x_max
1266
+ _, offset_y = fig.transFigure.inverted().transform((0, y_disp))
1267
+
1268
+ legend = fig.legend(
1269
+ handles=[legend_element],
1270
+ loc='center right',
1271
+ bbox_to_anchor=(offset_x, offset_y),
1272
+ borderaxespad=0.0,
1273
+ frameon=False,
1274
+ prop={'size': fontsize},
1275
+ )
1276
+ legend.get_frame().set_linewidth(0.5)
1277
+
1278
+ def _bootstrap_ci(series, estimator=np.median, ci=95, n_boot=1000, seed=42):
1279
+ rng = np.random.default_rng(seed)
1280
+ series = np.asarray(series)
1281
+ boot = np.array([estimator(
1282
+ resample(series, random_state=rng.integers(1e9))
1283
+ ) for _ in range(n_boot)])
1284
+ alpha = (100 - ci) / 2
1285
+ return (np.percentile(boot, alpha), np.percentile(boot, 100 - alpha))
1286
+
1287
+ def _legend_ci(data, fliersize, fontsize, ax, fig, order, ci=95, add_legend=True, x=None, y=None, hue=None):
1288
+ # add ci information
1289
+ x = x if x is not None else "configuration_id"
1290
+ y = y if y is not None else "quality"
1291
+ hue = hue if hue is not None else "configuration_id"
1292
+ grouped = data.groupby(x)[y]
1293
+
1294
+ medians = []
1295
+ lower = []
1296
+ upper = []
1297
+ # sort based on order
1298
+ for cid in order:
1299
+ s = grouped.get_group(cid)
1300
+
1301
+ m = np.median(s)
1302
+ lo, hi = _bootstrap_ci(s.to_numpy(), np.median, ci=ci)
1303
+
1304
+ medians.append(m)
1305
+ lower.append(lo)
1306
+ upper.append(hi)
1307
+
1308
+ medians = np.asarray(medians)
1309
+ lower = np.asarray(lower)
1310
+ upper = np.asarray(upper)
1311
+
1312
+ x_positions = np.arange(len(order)) - 0.25
1313
+ ax.errorbar(
1314
+ x=x_positions,
1315
+ y=medians,
1316
+ yerr=np.vstack([medians - lower, upper - medians]),
1317
+ fmt='o',
1318
+ color='red',
1319
+ capsize=4*fliersize,
1320
+ markersize=4*fliersize,
1321
+ elinewidth=2*fliersize,
1322
+ )
1323
+
1324
+ ax.set_xlabel("")
1325
+ ax.set_ylabel("")
1326
+
1327
+ if not add_legend: return
1328
+
1329
+ _add_legend(text='median (95% CI)', ax=ax, fig=fig, fontsize=fontsize)
1330
+
1331
+
1332
+ def _legend_rpd(data, fliersize, fontsize, ax, fig, order, add_legend=True, x=None, y=None, hue=None):
1333
+ # add rpd information
1334
+ x = x if x is not None else "configuration_id"
1335
+ y = y if y is not None else "quality"
1336
+ hue = hue if hue is not None else "configuration_id"
1337
+ grouped = data.groupby(x)[y]
1338
+
1339
+ medians = []
1340
+ lower = []
1341
+ upper = []
1342
+ # sort based on order
1343
+ for cid in order:
1344
+ s = grouped.get_group(cid)
1345
+
1346
+ m = s.median()
1347
+ mad = (s - m).abs().median()
1348
+
1349
+ medians.append(m)
1350
+ lower.append(m - mad)
1351
+ upper.append(m + mad)
1352
+
1353
+ medians = np.asarray(medians)
1354
+ lower = np.asarray(lower)
1355
+ upper = np.asarray(upper)
1356
+
1357
+ x_positions = np.arange(len(order)) - 0.25
1358
+
1359
+ ax.errorbar(
1360
+ x=x_positions,
1361
+ y=medians,
1362
+ yerr=np.vstack([medians - lower, upper - medians]),
1363
+ fmt='o',
1364
+ color='red',
1365
+ capsize=4*fliersize,
1366
+ markersize=4*fliersize,
1367
+ elinewidth=2*fliersize,
1368
+ )
1369
+
1370
+ ax.set_xlabel("")
1371
+ ax.set_ylabel("")
1372
+
1373
+ if not add_legend: return
1374
+
1375
+ _add_legend(text='median ± RPD', ax=ax, fig=fig, fontsize=fontsize)
1376
+
1377
+
1378
+ def _legend_ins(ax, ticks, labels):
1379
+ # add instance numbers
1380
+ dx = np.diff(ticks).mean()
1381
+ plt.xlim(ticks[0] - dx, ticks[-1] + dx)
1382
+
1383
+ t_top = ax.text(x=ticks[0], y=1.0, s="ins_num",
1384
+ ha='left', va='bottom',
1385
+ color='blue',
1386
+ transform=ax.get_xaxis_transform())
1387
+ t_bottom = []
1388
+ for i,x in enumerate(labels.values()):
1389
+ t = ax.text(x=ticks[i], y=0.99, s=x,
1390
+ ha='center', va='top',
1391
+ color='blue',
1392
+ transform=ax.get_xaxis_transform())
1393
+ t_bottom.append(t)
1394
+
1395
+ _align_left_to_text(ax, t_bottom[0], t_top)
1396
+
1397
+ def _align_left_to_text(ax, ref_text, target_text):
1398
+ fig = ax.figure
1399
+ fig.canvas.draw()
1400
+ renderer = fig.canvas.get_renderer()
1401
+
1402
+ bbox = ref_text.get_window_extent(renderer=renderer)
1403
+ x_left_disp = bbox.x0
1404
+
1405
+ x_left_data = ax.transData.inverted().transform((x_left_disp, 0))[0]
1406
+ target_text.set_x(x_left_data)