fix: 缺PyTorch时训练不会因ImportError中断整个流程

This commit is contained in:
tian 2026-06-10 09:41:19 +08:00
parent b6b072d844
commit 531cff509d

View File

@ -500,11 +500,15 @@ class ModelTrainer:
model.fit(X_train, y_train)
elif model_type == 'pytorch':
# 训练PyTorch模型
train_dataset = EquipmentDataset(X_train, y_train)
train_loader = DataLoader(train_dataset, batch_size=32, shuffle=True)
training_success = self.train_model(train_loader)
if not training_success:
# 训练PyTorch模型如未安装则跳过
try:
train_dataset = EquipmentDataset(X_train, y_train)
train_loader = DataLoader(train_dataset, batch_size=32, shuffle=True)
training_success = self.train_model(train_loader)
if not training_success:
continue
except ImportError as e:
logger.warning(f"PyTorch not available, skipping: {e}")
continue
# 评估模型性能