MCPcopy Create free account
hub / github.com/COLA-Laboratory/TransOPT / load_data

Function load_data

tests/data_analysis.py:12–34  ·  view source on GitHub ↗
(data_folder)

Source from the content-addressed store, hash-verified

10from mpl_toolkits.mplot3d import Axes3D
11
12def 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
36def get_non_dominated_solutions(data):
37 F = np.array([[1 - d['test_standard_acc'], 1 - d['test_robust_acc']] for d in data.values()])

Callers 4

compare_nsga2_resultsFunction · 0.70
data_analysis.pyFile · 0.70

Calls 1

loadMethod · 0.45

Tested by

no test coverage detected