From 35567e97835884b076a7c8e2006d1f59e91065b4 Mon Sep 17 00:00:00 2001
From: tian <11429339@qq.com>
Date: Wed, 10 Jun 2026 09:58:15 +0800
Subject: [PATCH] =?UTF-8?q?feat:=20=E8=AE=AD=E7=BB=83=E6=94=AF=E6=8C=81?=
=?UTF-8?q?=E8=87=AA=E5=8A=A8=E5=88=92=E5=88=86=E9=AA=8C=E8=AF=81=E9=9B=86?=
=?UTF-8?q?=EF=BC=8C=E6=97=A0=E9=9C=80=E6=89=8B=E5=B7=A5=E9=80=89=E6=8B=A9?=
=?UTF-8?q?=E4=B8=A4=E4=B8=AA=E6=95=B0=E6=8D=AE=E9=9B=86?=
MIME-Version: 1.0
Content-Type: text/plain; charset=UTF-8
Content-Transfer-Encoding: 8bit
---
frontend/src/views/TrainingPage.vue | 18 +++++++++++---
src/data_preparation.py | 30 +++++++++++++++++------
src/routes.py | 38 +++++++++++++++++++----------
3 files changed, 61 insertions(+), 25 deletions(-)
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()