1+ import argparse
12import os
23import sys
34import traceback
78
89try :
910 os .environ ['TF_CPP_MIN_LOG_LEVEL' ] = '2'
10-
1111 import tensorflow as tf
1212
1313 from cnsd import Dataset
1414 from cnsd .diagnosis .system import CNSD
1515 from cnsd .perception .cnn import _train_cnn
16- from cnsd .physics import PhysicsConfig
17- from validate_pu import load_pu_domain_split
1816
19- print ( 'Loading Authentic PU dataset (Cross-Domain RPM Split)...' )
20- ( X_train_full , y_train_full , cond_train_full ), ( X_target , y_target , cond_target ) = (
21- load_pu_domain_split ( )
22- )
17+ # Argparse
18+ parser = argparse . ArgumentParser ()
19+ parser . add_argument ( '--dataset' , type = str , required = True , choices = [ 'cwru' , 'pu' , 'xjtusy' ] )
20+ args = parser . parse_args ( )
2321
24- unique_rpm = set (cond_train_full ).union (set (cond_target ))
25- rpm_map = {float (r ): float (r ) for r in unique_rpm }
26- pu_physics = PhysicsConfig (
27- bearing = {'n_balls' : 8 , 'd_ball' : 6.75 , 'd_pitch' : 28.5 , 'contact_angle' : 0.0 },
28- cond_to_rpm = rpm_map ,
29- fs = 64000 ,
30- name = 'PU-6203' ,
31- )
32- pu_taxonomy = {
33- 0 : ('Normal' , 'None' ),
34- 1 : ('Outer Race' , 'Medium' ),
35- 2 : ('Inner Race' , 'High' ),
36- }
22+ print (f'Loading { args .dataset .upper ()} dataset...' )
23+
24+ if args .dataset == 'pu' :
25+ from cnsd .physics import PhysicsConfig
26+ from validate_pu import load_pu_domain_split
27+
28+ (X_train_full , y_train_full , cond_train_full ), (X_target , y_target , cond_target ) = (
29+ load_pu_domain_split ()
30+ )
31+ unique_rpm = set (cond_train_full ).union (set (cond_target ))
32+ rpm_map = {float (r ): float (r ) for r in unique_rpm }
33+ physics = PhysicsConfig (
34+ bearing = {'n_balls' : 8 , 'd_ball' : 6.75 , 'd_pitch' : 28.5 , 'contact_angle' : 0.0 },
35+ cond_to_rpm = rpm_map ,
36+ fs = 64000 ,
37+ name = 'PU-6203' ,
38+ )
39+ taxonomy = {0 : ('Normal' , 'None' ), 1 : ('Outer Race' , 'Medium' ), 2 : ('Inner Race' , 'High' )}
40+ fs = 64000
41+
42+ elif args .dataset == 'cwru' :
43+ from validate_run import CWRU , TAXONOMY , load_cwru
44+
45+ X , y , cond = load_cwru ()
46+ X = np .asarray (X , np .float32 )
47+ y = np .asarray (y )
48+ cond = np .asarray (cond )
49+
50+ train_mask = cond < 3
51+ target_mask = cond == 3
52+
53+ X_train_full = X [train_mask ]
54+ y_train_full = y [train_mask ]
55+ cond_train_full = cond [train_mask ]
56+
57+ X_target = X [target_mask ]
58+ y_target = y [target_mask ]
59+ cond_target = cond [target_mask ]
60+
61+ physics = CWRU
62+ taxonomy = TAXONOMY
63+ fs = 12000
64+
65+ elif args .dataset == 'xjtusy' :
66+ from cnsd .datasets .xjtusy import load_xjtusy_domain_split
67+
68+ train_ds , target_ds = load_xjtusy_domain_split (window_size = 4096 )
69+
70+ X_train_full = train_ds .X
71+ y_train_full = train_ds .y
72+ cond_train_full = train_ds .cond
73+
74+ X_target = target_ds .X
75+ y_target = target_ds .y
76+ cond_target = target_ds .cond
77+
78+ physics = target_ds .physics
79+ taxonomy = target_ds .taxonomy
80+ fs = target_ds .fs
3781
3882 def get_matched_coverage_gap (score , correct , target_n ):
3983 if target_n == 0 :
4084 return float ('nan' )
4185 # Score is higher for MORE confident
42- # Sort descending by score
4386 sorted_indices = np .argsort (score )[::- 1 ]
4487 hi_indices = sorted_indices [:target_n ]
4588
@@ -76,10 +119,10 @@ def clone_for_mc(layer):
76119 X_target ,
77120 y_target ,
78121 cond_target ,
79- fs = 64000 ,
80- physics = pu_physics ,
81- taxonomy = pu_taxonomy ,
82- name = 'PU_Test ' ,
122+ fs = fs ,
123+ physics = physics ,
124+ taxonomy = taxonomy ,
125+ name = f' { args . dataset . upper () } _Test ' ,
83126 )
84127 sig_te = np .stack ([test_ds .X [i ].reshape (- 1 ) for i in range (len (test_ds .X ))]).astype (np .float32 )
85128 yte = test_ds .y
@@ -112,10 +155,22 @@ def clone_for_mc(layer):
112155 )
113156
114157 train_ds = Dataset .from_arrays (
115- X_tr , y_tr , cond_tr , fs = 64000 , physics = pu_physics , taxonomy = pu_taxonomy , name = 'PU_Train'
158+ X_tr ,
159+ y_tr ,
160+ cond_tr ,
161+ fs = fs ,
162+ physics = physics ,
163+ taxonomy = taxonomy ,
164+ name = f'{ args .dataset .upper ()} _Train' ,
116165 )
117166 calib_ds = Dataset .from_arrays (
118- X_ca , y_ca , cond_ca , fs = 64000 , physics = pu_physics , taxonomy = pu_taxonomy , name = 'PU_Calib'
167+ X_ca ,
168+ y_ca ,
169+ cond_ca ,
170+ fs = fs ,
171+ physics = physics ,
172+ taxonomy = taxonomy ,
173+ name = f'{ args .dataset .upper ()} _Calib' ,
119174 )
120175
121176 # 2. Train Primary Model (Bypass SCM to prevent multiprocess deadlocks)
@@ -209,13 +264,10 @@ def clone_for_mc(layer):
209264 results ['mc_gap' ].append (mc_gap )
210265
211266 # Ensemble at matched coverage
212- ens_preds_probs = np .stack (
213- [m .predict (Xin_te , batch_size = 128 , verbose = 0 ) for m in ens ]
214- ) # (3, n, c)
215- ens_mean = ens_preds_probs .mean (0 ) # (n, c)
267+ ens_preds_probs = np .stack ([m .predict (Xin_te , batch_size = 128 , verbose = 0 ) for m in ens ])
268+ ens_mean = ens_preds_probs .mean (0 )
216269 ens_pred_class = ens_mean .argmax (1 )
217270 ens_correct = ens_pred_class == yte
218- # Score = negative entropy of ensemble mean
219271 ens_score = (ens_mean * np .log (ens_mean + eps )).sum (1 )
220272 ens_gap = get_matched_coverage_gap (ens_score , ens_correct , target_n )
221273 results ['ens_gap' ].append (ens_gap )
@@ -235,7 +287,6 @@ def clone_for_mc(layer):
235287 sig_n = sig_te + rng .randn (* sig_te .shape ).astype (np .float32 ) * np .sqrt (npow )
236288
237289 Xin_n = sig_n [..., None ]
238- # Use ensemble mode vote to define 'unanimous' exactly like Abhi's template
239290 ep_class = np .stack (
240291 [m .predict (Xin_n , batch_size = 128 , verbose = 0 ).argmax (1 ) for m in ens ]
241292 )
0 commit comments