Skip to content

Commit abc4f33

Browse files
committed
feat: Integrate XJTU-SY dataset and validate cross-domain domain adaptation
1 parent 5a35935 commit abc4f33

3 files changed

Lines changed: 226 additions & 0 deletions

File tree

cnsd/datasets/xjtusy.py

Lines changed: 117 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,117 @@
1+
import os
2+
import glob
3+
import numpy as np
4+
import pandas as pd
5+
from cnsd.datasets.contract import Dataset
6+
from cnsd.physics.configs import XJTUSY_PHYSICS
7+
8+
# Known fault mappings per XJTU-SY paper
9+
# We map Bearing ID to CNSD fault class: 1=Outer, 2=Inner, 3=Cage
10+
XJTUSY_FAULTS = {
11+
'35Hz12kN': {
12+
'Bearing1_1': 1, # Outer
13+
'Bearing1_2': 1, # Outer
14+
'Bearing1_3': 1, # Outer
15+
},
16+
'37.5Hz11kN': {
17+
'Bearing2_1': 2, # Inner
18+
'Bearing2_2': 1, # Outer
19+
'Bearing2_4': 1, # Outer
20+
'Bearing2_5': 1, # Outer
21+
},
22+
'40Hz10kN': {
23+
'Bearing3_1': 1, # Outer
24+
'Bearing3_2': 2, # Inner
25+
'Bearing3_3': 2, # Inner
26+
'Bearing3_4': 2, # Inner
27+
'Bearing3_5': 1, # Outer
28+
}
29+
}
30+
31+
def load_xjtusy_domain_split(data_dir=r'E:\301\CNSD\data\XJTU-SY\XJTU-SY_Bearing_Datasets', window_size=32768, train_cond='35Hz12kN', test_cond='37.5Hz11kN'):
32+
"""
33+
Loads authentic XJTU-SY dataset and strictly splits by condition (Domain Shift).
34+
Uses the run-to-failure nature to grab the first 15% as Healthy (0) and the
35+
last 15% as the Fault label.
36+
"""
37+
def _load_condition(cond_folder, rpm_val):
38+
X, y, cond = [], [], []
39+
cond_path = os.path.join(data_dir, cond_folder)
40+
41+
if not os.path.exists(cond_path):
42+
return [], [], []
43+
44+
for bearing_folder in os.listdir(cond_path):
45+
bearing_path = os.path.join(cond_path, bearing_folder)
46+
if not os.path.isdir(bearing_path):
47+
continue
48+
49+
fault_label = XJTUSY_FAULTS.get(cond_folder, {}).get(bearing_folder, None)
50+
if fault_label is None or bearing_folder == 'Bearing1_5': # skip 1_5 to avoid mixed labels
51+
continue
52+
53+
csv_files = sorted(glob.glob(os.path.join(bearing_path, '*.csv')),
54+
key=lambda x: int(os.path.splitext(os.path.basename(x))[0]))
55+
56+
total_files = len(csv_files)
57+
if total_files < 10:
58+
continue # ignore wildly corrupted directories
59+
60+
healthy_count = max(1, int(total_files * 0.20))
61+
fault_count = max(1, int(total_files * 0.20))
62+
63+
# Sliding window parameters
64+
step_size = 1024
65+
66+
# Extract Healthy
67+
for fpath in csv_files[:healthy_count]:
68+
df = pd.read_csv(fpath)
69+
sig = df.iloc[:, 0].values # Horizontal acceleration
70+
sig = (sig - np.mean(sig)) / (np.std(sig) + 1e-8)
71+
72+
# Slicing the 32768 array into overlapping 4096 windows
73+
for start_idx in range(0, len(sig) - window_size + 1, step_size):
74+
X.append(sig[start_idx : start_idx + window_size])
75+
y.append(0)
76+
cond.append(rpm_val)
77+
78+
# Extract Fault
79+
for fpath in csv_files[-fault_count:]:
80+
df = pd.read_csv(fpath)
81+
sig = df.iloc[:, 0].values
82+
sig = (sig - np.mean(sig)) / (np.std(sig) + 1e-8)
83+
84+
for start_idx in range(0, len(sig) - window_size + 1, step_size):
85+
X.append(sig[start_idx : start_idx + window_size])
86+
y.append(fault_label)
87+
cond.append(rpm_val)
88+
89+
return X, y, cond
90+
91+
X_train, y_train, c_train = _load_condition(train_cond, 2100.0)
92+
X_test, y_test, c_test = _load_condition(test_cond, 2250.0)
93+
94+
if not X_train or not X_test:
95+
raise FileNotFoundError(f"Could not load data. Ensure {data_dir} contains extracted {train_cond} and {test_cond} folders.")
96+
97+
ds_train = Dataset.from_arrays(
98+
X=np.array(X_train, dtype=np.float32),
99+
y=np.array(y_train, dtype=np.int32),
100+
cond=np.array(c_train, dtype=np.float32),
101+
fs=25600,
102+
physics=XJTUSY_PHYSICS,
103+
taxonomy={0: ('Normal', 'None'), 1: ('Outer Race', 'Medium'), 2: ('Inner Race', 'High')},
104+
name='XJTUSY_Train'
105+
)
106+
107+
ds_test = Dataset.from_arrays(
108+
X=np.array(X_test, dtype=np.float32),
109+
y=np.array(y_test, dtype=np.int32),
110+
cond=np.array(c_test, dtype=np.float32),
111+
fs=25600,
112+
physics=XJTUSY_PHYSICS,
113+
taxonomy={0: ('Normal', 'None'), 1: ('Outer Race', 'Medium'), 2: ('Inner Race', 'High')},
114+
name='XJTUSY_Test'
115+
)
116+
117+
return ds_train, ds_test

cnsd/physics/configs.py

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -31,3 +31,10 @@ class PhysicsConfig:
3131
fs=20000,
3232
name='SEU-Gearbox',
3333
)
34+
35+
XJTUSY_PHYSICS = PhysicsConfig(
36+
bearing={'n_balls': 8, 'd_ball': 7.94, 'd_pitch': 34.55, 'contact_angle': 0.0},
37+
cond_to_rpm={2100.0: 2100.0, 2250.0: 2250.0, 2400.0: 2400.0},
38+
fs=25600,
39+
name='XJTU-SY-LDK-UER204',
40+
)

validate_xjtusy.py

Lines changed: 102 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,102 @@
1+
import numpy as np
2+
import tensorflow as tf
3+
4+
from cnsd import Dataset
5+
from cnsd.diagnosis.system import CNSD
6+
from cnsd.datasets.xjtusy import load_xjtusy_domain_split
7+
8+
def headline_accuracy_by_verdict(report, y_true):
9+
pred = np.array([r['predicted_class'] for r in report.records])
10+
correct = pred == np.asarray(y_true)
11+
verdicts = np.array([r['physics_verdict'] for r in report.records])
12+
out = {}
13+
for v in ('CONFIRMED', 'CONFLICT', 'INCONCLUSIVE'):
14+
m = verdicts == v
15+
if m.any():
16+
out[v] = {'n': int(m.sum()), 'cnn_accuracy': float(correct[m].mean())}
17+
return out
18+
19+
if __name__ == '__main__':
20+
np.random.seed(42)
21+
tf.random.set_seed(42)
22+
23+
print('Loading XJTU-SY dataset (Cross-Domain RPM/Load Split)...')
24+
train_data, target_data = load_xjtusy_domain_split(window_size=4096)
25+
26+
# Split target domain into Calib (50%) and Test (50%)
27+
indices = np.arange(len(target_data.y))
28+
np.random.shuffle(indices)
29+
calib_size = len(indices) // 2
30+
31+
calib_idx = indices[:calib_size]
32+
test_idx = indices[calib_size:]
33+
34+
X_calib, y_calib, cond_calib = target_data.X[calib_idx], target_data.y[calib_idx], target_data.cond[calib_idx]
35+
X_test, y_test, cond_test = target_data.X[test_idx], target_data.y[test_idx], target_data.cond[test_idx]
36+
37+
calib_data = Dataset.from_arrays(
38+
X_calib, y_calib, cond_calib,
39+
fs=target_data.fs, physics=target_data.physics,
40+
taxonomy=target_data.taxonomy,
41+
name='XJTUSY_Calib'
42+
)
43+
test_data = Dataset.from_arrays(
44+
X_test, y_test, cond_test,
45+
fs=target_data.fs, physics=target_data.physics,
46+
taxonomy=target_data.taxonomy,
47+
name='XJTUSY_Test'
48+
)
49+
50+
print(f'Train (2100 RPM)={len(train_data.y)} | Calib (2250 RPM)={len(y_calib)} | Test (2250 RPM)={len(y_test)}')
51+
52+
model = CNSD()
53+
54+
print('\n[1] Training Neural Network on 2100 RPM Source Data...')
55+
model.fit(train_data, epochs=20)
56+
57+
print('\n[2] Calibrating Tau threshold on 2250 RPM Target Data...')
58+
taus = np.arange(1.0, 5.1, 0.5)
59+
best_gap = -np.inf
60+
best_tau = 1.0
61+
62+
for tau in taus:
63+
model.symbolic.tau = float(tau)
64+
report = model.diagnose(calib_data)
65+
66+
hb = headline_accuracy_by_verdict(report, y_calib)
67+
68+
conf_acc = hb.get('CONFIRMED', {}).get('cnn_accuracy', 0.0)
69+
cnfl_acc = hb.get('CONFLICT', {}).get('cnn_accuracy', 0.0)
70+
gap = conf_acc - cnfl_acc if 'CONFIRMED' in hb and 'CONFLICT' in hb else 0.0
71+
72+
print(f'Calib tau={tau:.1f} | Conf={conf_acc:.3f} | Cnfl={cnfl_acc:.3f} | Gap={gap:+.3f}')
73+
if gap > best_gap:
74+
best_gap = gap
75+
best_tau = float(tau)
76+
77+
print(f'\n=> Selected optimal tau: {best_tau}')
78+
79+
print('\n[3] Evaluating on Test Set (2250 RPM)...')
80+
model.symbolic.tau = best_tau
81+
report = model.diagnose(test_data)
82+
pred = np.array([r['predicted_class'] for r in report.records])
83+
baseline_acc = float((pred == np.asarray(y_test)).mean())
84+
85+
print('\n--- FINAL TEST RESULTS (CROSS-DOMAIN XJTU-SY) ---')
86+
print(f'Baseline CNN Acc: {baseline_acc:.3f}')
87+
print('--------------------------------------------')
88+
89+
hb = headline_accuracy_by_verdict(report, y_test)
90+
if 'CONFIRMED' in hb:
91+
print(f' Physics-Confirmed Acc: {hb["CONFIRMED"]["cnn_accuracy"]:.3f} (n={hb["CONFIRMED"]["n"]})')
92+
if 'CONFLICT' in hb:
93+
print(f' Physics-Conflict Acc: {hb["CONFLICT"]["cnn_accuracy"]:.3f} (n={hb["CONFLICT"]["n"]})')
94+
if 'INCONCLUSIVE' in hb:
95+
inc_n = hb['INCONCLUSIVE']['n']
96+
inc_pct = (inc_n / len(y_test)) * 100
97+
print(f' Physics-Inconclusive Acc:{hb["INCONCLUSIVE"]["cnn_accuracy"]:.3f} (n={inc_n}, {inc_pct:.1f}%)')
98+
99+
if 'CONFIRMED' in hb and 'CONFLICT' in hb:
100+
gap = hb['CONFIRMED']['cnn_accuracy'] - hb['CONFLICT']['cnn_accuracy']
101+
print(f' GAP (CONF - CNFL): {gap:+.3f}')
102+
print('--------------------------------------------')

0 commit comments

Comments
 (0)