|
6 | 6 | import numpy as np |
7 | 7 | import scipy.stats as stats |
8 | 8 |
|
9 | | -try: |
| 9 | + |
| 10 | +def main(): |
10 | 11 | os.environ['TF_CPP_MIN_LOG_LEVEL'] = '2' |
11 | 12 | import tensorflow as tf |
12 | 13 |
|
13 | 14 | from cnsd import Dataset |
14 | 15 | from cnsd.diagnosis.system import CNSD |
15 | 16 | from cnsd.perception.cnn import _train_cnn |
16 | 17 |
|
17 | | - # Argparse |
18 | 18 | parser = argparse.ArgumentParser() |
19 | 19 | parser.add_argument('--dataset', type=str, required=True, choices=['cwru', 'pu', 'xjtusy']) |
20 | 20 | args = parser.parse_args() |
|
23 | 23 |
|
24 | 24 | if args.dataset == 'pu': |
25 | 25 | from cnsd.physics import PhysicsConfig |
26 | | - from validate_pu import load_pu_domain_split |
| 26 | + from validation.validate_pu import load_pu_domain_split |
27 | 27 |
|
28 | 28 | (X_train_full, y_train_full, cond_train_full), (X_target, y_target, cond_target) = ( |
29 | 29 | load_pu_domain_split() |
|
40 | 40 | fs = 64000 |
41 | 41 |
|
42 | 42 | elif args.dataset == 'cwru': |
43 | | - from validate_run import CWRU, TAXONOMY, load_cwru |
| 43 | + from validation.validate_cwru import CWRU, TAXONOMY, load_cwru |
44 | 44 |
|
45 | 45 | X, y, cond = load_cwru() |
46 | 46 | X = np.asarray(X, np.float32) |
@@ -339,7 +339,11 @@ def clone_for_mc(layer): |
339 | 339 | print('================ DONE ================') |
340 | 340 |
|
341 | 341 | except Exception: |
342 | | - with open('crash_traceback.txt', 'w') as f: |
343 | | - traceback.print_exc(file=f) |
344 | | - print('CRASHED. Check crash_traceback.txt') |
345 | | - sys.exit(1) |
| 342 | + with open('crash_traceback.txt', 'w') as f: |
| 343 | + traceback.print_exc(file=f) |
| 344 | + print('CRASHED. Check crash_traceback.txt') |
| 345 | + sys.exit(1) |
| 346 | + |
| 347 | + |
| 348 | +if __name__ == '__main__': |
| 349 | + main() |
0 commit comments