feat: 训练支持自动划分验证集,无需手工选择两个数据集
This commit is contained in:
parent
67b370364e
commit
35567e9783
@ -16,7 +16,7 @@
|
|||||||
</el-form-item>
|
</el-form-item>
|
||||||
|
|
||||||
<el-form-item label="训练数据集">
|
<el-form-item label="训练数据集">
|
||||||
<el-select v-model="trainingConfig.train_dataset_id" placeholder="选择训练数据集">
|
<el-select v-model="trainingConfig.train_dataset_id" placeholder="选择数据集">
|
||||||
<el-option
|
<el-option
|
||||||
v-for="dataset in trainingDatasets"
|
v-for="dataset in trainingDatasets"
|
||||||
:key="dataset.id"
|
:key="dataset.id"
|
||||||
@ -26,7 +26,15 @@
|
|||||||
</el-select>
|
</el-select>
|
||||||
</el-form-item>
|
</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-select v-model="trainingConfig.validation_dataset_id" placeholder="选择验证数据集">
|
||||||
<el-option
|
<el-option
|
||||||
v-for="dataset in validationDatasets"
|
v-for="dataset in validationDatasets"
|
||||||
@ -299,6 +307,8 @@ const trainingConfig = ref({
|
|||||||
type: '',
|
type: '',
|
||||||
train_dataset_id: null,
|
train_dataset_id: null,
|
||||||
validation_dataset_id: null,
|
validation_dataset_id: null,
|
||||||
|
auto_split: true,
|
||||||
|
split_ratio: 0.2,
|
||||||
models: ['pytorch', 'xgboost', 'lightgbm', 'gbm', 'rf']
|
models: ['pytorch', 'xgboost', 'lightgbm', 'gbm', 'rf']
|
||||||
})
|
})
|
||||||
|
|
||||||
@ -352,8 +362,8 @@ const startTraining = async () => {
|
|||||||
ElMessage.warning('请选择训练数据集')
|
ElMessage.warning('请选择训练数据集')
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
if (!trainingConfig.value.validation_dataset_id) {
|
if (!trainingConfig.value.auto_split && !trainingConfig.value.validation_dataset_id) {
|
||||||
ElMessage.warning('请选择验证数据集')
|
ElMessage.warning('请选择验证数据集或开启自动划分')
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
if (trainingConfig.value.models.length === 0) {
|
if (trainingConfig.value.models.length === 0) {
|
||||||
|
|||||||
@ -40,8 +40,8 @@ class DataPreparation:
|
|||||||
self.feature_scaler = StandardScaler()
|
self.feature_scaler = StandardScaler()
|
||||||
self.target_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:
|
try:
|
||||||
logger.info(f"Preparing training data for {equipment_type}")
|
logger.info(f"Preparing training data for {equipment_type}")
|
||||||
|
|
||||||
@ -91,12 +91,7 @@ class DataPreparation:
|
|||||||
X_scaled = self.feature_scaler.fit_transform(X)
|
X_scaled = self.feature_scaler.fit_transform(X)
|
||||||
y_scaled = self.target_scaler.fit_transform(y.reshape(-1, 1)).ravel()
|
y_scaled = self.target_scaler.fit_transform(y.reshape(-1, 1)).ravel()
|
||||||
|
|
||||||
# 创建数据集和数据加载器
|
result = {
|
||||||
dataset = EquipmentDataset(X_scaled, y_scaled)
|
|
||||||
dataloader = DataLoader(dataset, batch_size=batch_size, shuffle=True)
|
|
||||||
|
|
||||||
return {
|
|
||||||
'dataloader': dataloader,
|
|
||||||
'feature_names': feature_names,
|
'feature_names': feature_names,
|
||||||
'feature_scaler': self.feature_scaler,
|
'feature_scaler': self.feature_scaler,
|
||||||
'target_scaler': self.target_scaler,
|
'target_scaler': self.target_scaler,
|
||||||
@ -105,6 +100,25 @@ class DataPreparation:
|
|||||||
'y': y_scaled
|
'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:
|
except Exception as e:
|
||||||
logger.error(f"Error in data preparation: {str(e)}")
|
logger.error(f"Error in data preparation: {str(e)}")
|
||||||
raise Exception(f"Training error: {str(e)}")
|
raise Exception(f"Training error: {str(e)}")
|
||||||
|
|||||||
@ -303,6 +303,8 @@ def train_model():
|
|||||||
train_dataset_id = data.get('train_dataset_id')
|
train_dataset_id = data.get('train_dataset_id')
|
||||||
validation_dataset_id = data.get('validation_dataset_id')
|
validation_dataset_id = data.get('validation_dataset_id')
|
||||||
models = data.get('models', [])
|
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"Training dataset: {train_dataset_id}")
|
||||||
logger.info(f"Validation dataset: {validation_dataset_id}")
|
logger.info(f"Validation dataset: {validation_dataset_id}")
|
||||||
@ -383,9 +385,19 @@ def train_model():
|
|||||||
data_processor = DataPreparation()
|
data_processor = DataPreparation()
|
||||||
|
|
||||||
# 准备训练数据
|
# 准备训练数据
|
||||||
|
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)
|
train_prepared = data_processor.prepare_training_data(train_data, equipment_type)
|
||||||
|
|
||||||
# 准备验证数据(如果有)
|
|
||||||
validation_prepared = None
|
validation_prepared = None
|
||||||
if validation_data:
|
if validation_data:
|
||||||
validation_prepared = data_processor.prepare_validation_data(
|
validation_prepared = data_processor.prepare_validation_data(
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user