diff --git a/frontend/src/views/TrainingPage.vue b/frontend/src/views/TrainingPage.vue index 0393cf6..06bb152 100644 --- a/frontend/src/views/TrainingPage.vue +++ b/frontend/src/views/TrainingPage.vue @@ -16,7 +16,7 @@ - + - + + + + + + + + + { ElMessage.warning('请选择训练数据集') return } - if (!trainingConfig.value.validation_dataset_id) { - ElMessage.warning('请选择验证数据集') + if (!trainingConfig.value.auto_split && !trainingConfig.value.validation_dataset_id) { + ElMessage.warning('请选择验证数据集或开启自动划分') return } if (trainingConfig.value.models.length === 0) { diff --git a/src/data_preparation.py b/src/data_preparation.py index 48cb5f7..bec1664 100644 --- a/src/data_preparation.py +++ b/src/data_preparation.py @@ -40,8 +40,8 @@ class DataPreparation: self.feature_scaler = StandardScaler() self.target_scaler = StandardScaler() - def prepare_training_data(self, equipment_data, equipment_type, batch_size=32): - """准备训练数据""" + def prepare_training_data(self, equipment_data, equipment_type, batch_size=32, validation_split=None): + """准备训练数据,可选自动划分验证集""" try: logger.info(f"Preparing training data for {equipment_type}") @@ -91,12 +91,7 @@ class DataPreparation: X_scaled = self.feature_scaler.fit_transform(X) y_scaled = self.target_scaler.fit_transform(y.reshape(-1, 1)).ravel() - # 创建数据集和数据加载器 - dataset = EquipmentDataset(X_scaled, y_scaled) - dataloader = DataLoader(dataset, batch_size=batch_size, shuffle=True) - - return { - 'dataloader': dataloader, + result = { 'feature_names': feature_names, 'feature_scaler': self.feature_scaler, 'target_scaler': self.target_scaler, @@ -105,6 +100,25 @@ class DataPreparation: 'y': y_scaled } + # 如果指定了验证集比例,自动划分 + if validation_split is not None and 0 < validation_split < 1: + from sklearn.model_selection import train_test_split + X_train, X_val, y_train, y_val = train_test_split( + X_scaled, y_scaled, test_size=validation_split, random_state=42 + ) + result['X'] = X_train + result['y'] = y_train + result['X_val'] = X_val + result['y_val'] = y_val + logger.info(f"Auto-split: train={X_train.shape[0]}, val={X_val.shape[0]}") + + # 创建数据集和数据加载器(仅用于训练集) + dataset = EquipmentDataset(result['X'], result['y']) + dataloader = DataLoader(dataset, batch_size=batch_size, shuffle=True) + result['dataloader'] = dataloader + + return result + except Exception as e: logger.error(f"Error in data preparation: {str(e)}") raise Exception(f"Training error: {str(e)}") diff --git a/src/routes.py b/src/routes.py index 9e9864d..3823b83 100644 --- a/src/routes.py +++ b/src/routes.py @@ -303,6 +303,8 @@ def train_model(): train_dataset_id = data.get('train_dataset_id') validation_dataset_id = data.get('validation_dataset_id') models = data.get('models', []) + auto_split = data.get('auto_split', False) # 自动划分训练/验证集 + split_ratio = data.get('split_ratio', 0.2) # 验证集比例,默认 20% logger.info(f"Training dataset: {train_dataset_id}") logger.info(f"Validation dataset: {validation_dataset_id}") @@ -383,20 +385,30 @@ def train_model(): data_processor = DataPreparation() # 准备训练数据 - train_prepared = data_processor.prepare_training_data(train_data, equipment_type) - - # 准备验证数据(如果有) - validation_prepared = None - if validation_data: - validation_prepared = data_processor.prepare_validation_data( - validation_data, - equipment_type, - train_prepared['feature_names'], - { - 'feature_scaler': train_prepared['feature_scaler'], - 'target_scaler': train_prepared['target_scaler'] - } + if auto_split and not validation_dataset_id: + # 自动划分模式:从训练数据中随机拆分验证集 + logger.info(f"Auto-splitting training data with validation ratio={split_ratio}") + train_prepared = data_processor.prepare_training_data( + train_data, equipment_type, validation_split=split_ratio ) + validation_prepared = { + 'X': train_prepared.pop('X_val'), + 'y': train_prepared.pop('y_val'), + } + else: + # 传统模式:单独查询验证数据集 + train_prepared = data_processor.prepare_training_data(train_data, equipment_type) + validation_prepared = None + if validation_data: + validation_prepared = data_processor.prepare_validation_data( + validation_data, + equipment_type, + train_prepared['feature_names'], + { + 'feature_scaler': train_prepared['feature_scaler'], + 'target_scaler': train_prepared['target_scaler'] + } + ) # 2. 训练模型 model_trainer = ModelTrainer()