feat: 训练支持自动划分验证集,无需手工选择两个数据集

This commit is contained in:
tian 2026-06-10 09:58:15 +08:00
parent 67b370364e
commit 35567e9783
3 changed files with 61 additions and 25 deletions

View File

@ -16,7 +16,7 @@
</el-form-item>
<el-form-item label="训练数据集">
<el-select v-model="trainingConfig.train_dataset_id" placeholder="选择训练数据集">
<el-select v-model="trainingConfig.train_dataset_id" placeholder="选择数据集">
<el-option
v-for="dataset in trainingDatasets"
:key="dataset.id"
@ -26,7 +26,15 @@
</el-select>
</el-form-item>
<el-form-item label="验证数据集">
<el-form-item label="验证集划分">
<el-switch v-model="trainingConfig.auto_split" active-text="自动" inactive-text="手动" />
</el-form-item>
<el-form-item v-if="trainingConfig.auto_split" label="验证集比例">
<el-slider v-model="trainingConfig.split_ratio" :min="0.1" :max="0.4" :step="0.05" show-input style="width: 300px" />
</el-form-item>
<el-form-item v-else label="验证数据集">
<el-select v-model="trainingConfig.validation_dataset_id" placeholder="选择验证数据集">
<el-option
v-for="dataset in validationDatasets"
@ -299,6 +307,8 @@ const trainingConfig = ref({
type: '',
train_dataset_id: null,
validation_dataset_id: null,
auto_split: true,
split_ratio: 0.2,
models: ['pytorch', 'xgboost', 'lightgbm', 'gbm', 'rf']
})
@ -352,8 +362,8 @@ const startTraining = async () => {
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) {

View File

@ -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)}")

View File

@ -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()