(data_folder)
| 10 | from mpl_toolkits.mplot3d import Axes3D |
| 11 | |
| 12 | def load_data(data_folder): |
| 13 | data = {} |
| 14 | for filename in os.listdir(data_folder): |
| 15 | file_path = os.path.join(data_folder, filename) |
| 16 | if os.path.isfile(file_path): |
| 17 | with open(file_path, 'r') as f: |
| 18 | content = json.load(f) |
| 19 | x = [] |
| 20 | for key in ['lr', 'weight_decay', 'momentum', 'dropout_rate']: |
| 21 | pattern = rf'({key}_)([\d.e-]+)' |
| 22 | match = re.search(pattern, filename) |
| 23 | if match: |
| 24 | value = float(match.group(2)) |
| 25 | if key in ['lr', 'weight_decay']: |
| 26 | x.append(np.log10(value)) |
| 27 | else: |
| 28 | x.append(value) |
| 29 | data[filename] = { |
| 30 | 'x': x, |
| 31 | 'test_standard_acc': content['test_standard_acc'], |
| 32 | 'test_robust_acc': np.mean([v for k, v in content.items() if k.startswith('test_') and k != 'test_standard_acc']) |
| 33 | } |
| 34 | return data |
| 35 | |
| 36 | def get_non_dominated_solutions(data): |
| 37 | F = np.array([[1 - d['test_standard_acc'], 1 - d['test_robust_acc']] for d in data.values()]) |
no test coverage detected