1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93
| from sklearn.ensemble import RandomForestClassifier, GradientBoostingClassifier from sklearn.svm import SVC from sklearn.model_selection import cross_val_score from sklearn.preprocessing import StandardScaler from sklearn.pipeline import Pipeline
class HRVDriverStateClassifier: """ HRV驾驶员状态分类器 支持三种分类器: 1. Random Forest 2. Gradient Boosting 3. SVM (RBF核) 二分类: Rested vs Tired """ def __init__(self, classifier="rf"): self.scaler = StandardScaler() if classifier == "rf": self.model = RandomForestClassifier( n_estimators=100, max_depth=8, random_state=42 ) elif classifier == "gb": self.model = GradientBoostingClassifier( n_estimators=100, max_depth=4, random_state=42 ) elif classifier == "svm": self.model = SVC(kernel="rbf", C=1.0, gamma="scale") self.pipeline = Pipeline([ ("scaler", self.scaler), ("classifier", self.model) ]) def extract_features(self, rr_intervals): """提取完整HRV特征集""" time_features = compute_time_domain_hrv(rr_intervals) freq_features = compute_frequency_domain_hrv(rr_intervals) return {**time_features, **freq_features} def train(self, rr_samples, labels): """ 训练分类器 Args: rr_samples: RR间期列表, 每个元素是一段RR序列 labels: 0=休息, 1=疲劳 """ feature_matrix = [] for rr in rr_samples: features = self.extract_features(rr) feature_matrix.append(list(features.values())) X = np.array(feature_matrix) y = np.array(labels) scores = cross_val_score(self.pipeline, X, y, cv=5, scoring="accuracy") print(f"交叉验证准确率: {scores.mean():.3f} ± {scores.std():.3f}") self.pipeline.fit(X, y) return scores
np.random.seed(42) n_samples = 200
rested_samples = [] for _ in range(n_samples // 2): rr = np.random.normal(0.85, np.random.uniform(0.06, 0.10), 1800) rested_samples.append(rr)
tired_samples = [] for _ in range(n_samples // 2): rr = np.random.normal(0.75, np.random.uniform(0.03, 0.05), 1800) tired_samples.append(rr)
all_samples = rested_samples + tired_samples all_labels = [0] * (n_samples // 2) + [1] * (n_samples // 2)
for clf_name in ["rf", "gb", "svm"]: print(f"\n=== {clf_name.upper()} 分类器 ===") classifier = HRVDriverStateClassifier(classifier=clf_name) scores = classifier.train(all_samples, all_labels)
|