From fccd4c43669535ce4dbd4e7edf4457c8f0f362d0 Mon Sep 17 00:00:00 2001
From: Tian jianyong <11429339@qq.com>
Date: Sat, 9 Nov 2024 16:48:50 +0800
Subject: [PATCH] =?UTF-8?q?=E5=AE=8C=E6=88=90=E4=BA=86=E5=9F=BA=E6=9C=AC?=
=?UTF-8?q?=E5=8A=9F=E8=83=BD?=
MIME-Version: 1.0
Content-Type: text/plain; charset=UTF-8
Content-Transfer-Encoding: 8bit
---
.cursorrules | 24 --
app.py | 8 -
docs/debug.md | 29 ++
frontend/jsconfig.json | 2 +-
frontend/src/App.vue | 7 +-
frontend/src/main.js | 4 +-
frontend/src/router/index.js | 6 -
frontend/src/views/AnalysisPage.vue | 3 -
frontend/src/views/DataPage.vue | 2 +-
frontend/src/views/HomePage.vue | 26 +-
frontend/src/views/ModelPage.vue | 21 +-
frontend/src/views/PLSPredictPage.vue | 243 -------------
frontend/src/views/PredictPage.vue | 237 ++++++++----
frontend/src/views/TrainingPage.vue | 374 +++++++++++--------
run.py | 66 +---
src/api.py | 2 +-
src/app.py | 96 ++---
src/cost_prediction.py | 64 +++-
src/create_template.py | 3 +
src/data_preparation.py | 85 +++--
src/database/db_connection.py | 37 +-
src/feature_analysis.py | 11 +-
src/import_data.py | 52 ++-
src/logger.py | 33 ++
src/model_trainer.py | 494 ++++++++++++++++++--------
src/pls_regression.py | 313 ----------------
src/routes.py | 366 +++++++++----------
src/run.py | 28 --
28 files changed, 1211 insertions(+), 1425 deletions(-)
delete mode 100644 app.py
delete mode 100644 frontend/src/views/PLSPredictPage.vue
create mode 100644 src/logger.py
delete mode 100644 src/pls_regression.py
delete mode 100644 src/run.py
diff --git a/.cursorrules b/.cursorrules
index f9d19c8..5815276 100644
--- a/.cursorrules
+++ b/.cursorrules
@@ -1,25 +1,3 @@
-# 开发流程
-
-First ensure basic functionality works
-Implement core functionality using the simplest direct approach
-Ensure data flow is working correctly
-Verify results are accurate
-Then gradually add additional features
-Add error handling
-Add data validation
-Add format conversion
-Add logging
-Improve user experience
-Avoid premature optimization
-Don't do complex data validation at the start
-Don't worry about performance optimization early
-Don't over-engineer
-This development flow:
-Quickly validates if core functionality works
-Identifies and fixes fundamental issues early
-Avoids wasting time on unnecessary optimizations
-Makes code easier to maintain and debug
-These principles should guide all code responses, focusing on getting the basics working first before adding complexity.
# 代码修改最佳实践
@@ -78,5 +56,3 @@ These principles should guide all code responses, focusing on getting the basics
- 处理异常情况
- 保护敏感信息
- 添加访问控制
-
-These practices help maintain code quality and reduce potential issues.
diff --git a/app.py b/app.py
deleted file mode 100644
index da4dbaf..0000000
--- a/app.py
+++ /dev/null
@@ -1,8 +0,0 @@
-import logging
-
-# 配置日志
-logging.basicConfig(
- filename='logs/api.log',
- level=logging.INFO,
- format='%(asctime)s - %(levelname)s - %(message)s'
-)
\ No newline at end of file
diff --git a/docs/debug.md b/docs/debug.md
index 0d78604..ce13431 100644
--- a/docs/debug.md
+++ b/docs/debug.md
@@ -614,3 +614,32 @@ trainingResult.value = null
2. 可以考虑集成 XGBoost 和 Random Forest
3. 继续调整 LightGBM 的参数
4. 暂时不使用 GBDT
+
+### 数据集存在的问题
+
+火箭炮数据集:
+
+- Feature length_m missing rate: 0.00%
+- Feature width_m missing rate: 9.09%
+- Feature height_m missing rate: 9.09%
+- Feature weight_kg missing rate: 0.00%
+- Feature max_range_km missing rate: 45.45%
+- Feature firing_angle_horizontal missing rate: 45.45%
+- Feature firing_angle_vertical missing rate: 45.45%
+- Feature rocket_length_m missing rate: 72.73%
+- Feature rocket_diameter_mm missing rate: 0.00%
+- Feature rocket_weight_kg missing rate: 72.73%
+- Feature rate_of_fire missing rate: 54.55%
+
+巡飞弹数据集:
+
+- Feature length_m missing rate: 27.78%
+- Feature width_m missing rate: 50.00%
+- Feature height_m missing rate: 50.00%
+- Feature weight_kg missing rate: 22.22%
+- Feature max_range_km missing rate: 44.44%
+- Feature wingspan_m missing rate: 50.00%
+- Feature warhead_weight_kg missing rate: 77.78%
+- Feature max_speed_ms missing rate: 77.78%
+- Feature cruise_speed_kmh missing rate: 61.11%
+- Feature flight_time_min missing rate: 33.33%
diff --git a/frontend/jsconfig.json b/frontend/jsconfig.json
index 2ee5342..303c328 100644
--- a/frontend/jsconfig.json
+++ b/frontend/jsconfig.json
@@ -16,5 +16,5 @@
"scripthost"
]
},
- "exclude": ["**/HelloWorld.vue"]
+ "include": ["src/**/*"]
}
diff --git a/frontend/src/App.vue b/frontend/src/App.vue
index 7afc899..c4278c5 100644
--- a/frontend/src/App.vue
+++ b/frontend/src/App.vue
@@ -7,13 +7,12 @@
:default-active="$route.path"
>
首页
- 机器学习预测
- PLS回归预测
+ 成本预测
特征分析
模型训练
- 数据管理
- 数据集管理
模型管理
+ 数据集管理
+ 数据管理
diff --git a/frontend/src/main.js b/frontend/src/main.js
index 7763e6f..360a4a4 100644
--- a/frontend/src/main.js
+++ b/frontend/src/main.js
@@ -23,7 +23,7 @@ for (const [key, component] of Object.entries(ElementPlusIconsVue)) {
}
// 全局错误处理
-app.config.errorHandler = (err, vm, info) => {
+app.config.errorHandler = (err) => {
if (err.message && err.message.includes('ResizeObserver')) {
return
}
@@ -31,7 +31,7 @@ app.config.errorHandler = (err, vm, info) => {
}
// 全局警告处理
-app.config.warnHandler = (msg, vm, trace) => {
+app.config.warnHandler = (msg, trace) => {
if (msg.includes('ResizeObserver')) {
return
}
diff --git a/frontend/src/router/index.js b/frontend/src/router/index.js
index ec88e5d..3dfcdee 100644
--- a/frontend/src/router/index.js
+++ b/frontend/src/router/index.js
@@ -3,7 +3,6 @@ import HomePage from '@/views/HomePage.vue'
import DataPage from '@/views/DataPage.vue'
import DatasetPage from '@/views/DatasetPage.vue'
import PredictPage from '@/views/PredictPage.vue'
-import PLSPredictPage from '@/views/PLSPredictPage.vue'
import AnalysisPage from '@/views/AnalysisPage.vue'
import TrainingPage from '@/views/TrainingPage.vue'
@@ -28,11 +27,6 @@ const routes = [
name: 'Predict',
component: PredictPage
},
- {
- path: '/pls-predict',
- name: 'PLSPredict',
- component: PLSPredictPage
- },
{
path: '/analysis',
name: 'Analysis',
diff --git a/frontend/src/views/AnalysisPage.vue b/frontend/src/views/AnalysisPage.vue
index 23bc86b..f8f8ba1 100644
--- a/frontend/src/views/AnalysisPage.vue
+++ b/frontend/src/views/AnalysisPage.vue
@@ -71,9 +71,6 @@ import axios from 'axios'
import { API_BASE_URL } from '@/config'
import * as echarts from 'echarts'
-// 定义组件名称
-const __name = 'AnalysisPage'
-
// 响应式数据
const analysisForm = ref({
equipment_type: '',
diff --git a/frontend/src/views/DataPage.vue b/frontend/src/views/DataPage.vue
index 3c72c67..4689f70 100644
--- a/frontend/src/views/DataPage.vue
+++ b/frontend/src/views/DataPage.vue
@@ -647,7 +647,7 @@ const isNumericParam = (param) => {
}
// 添加 handleTabClick 函数
-const handleTabClick = (tab) => {
+const handleTabClick = () => {
// 切换标签页时重置过滤条件
searchQuery.value = ''
filterManufacturer.value = ''
diff --git a/frontend/src/views/HomePage.vue b/frontend/src/views/HomePage.vue
index 4039d20..60c0676 100644
--- a/frontend/src/views/HomePage.vue
+++ b/frontend/src/views/HomePage.vue
@@ -8,15 +8,8 @@
- 机器学习预测
- 基于机器学习模型的成本预测
-
-
-
-
-
- PLS回归预测
- 基于偏最小二乘回归的成本预测
+ 成本预测
+ 基于机器学习和 PLS 回归模型的成本预测
@@ -34,10 +27,10 @@
-
+
- 数据管理
- 管理装备数据和成本数据
+ 模型管理
+ 管理训练好的模型
@@ -47,13 +40,20 @@
管理训练和验证数据集
+
+
+
+ 数据管理
+ 管理装备数据和成本数据
+
+
\ No newline at end of file
diff --git a/frontend/src/views/PredictPage.vue b/frontend/src/views/PredictPage.vue
index 7001002..a42fa9f 100644
--- a/frontend/src/views/PredictPage.vue
+++ b/frontend/src/views/PredictPage.vue
@@ -2,11 +2,11 @@
- 装备成本预测
+ 成本预测
-
+
@@ -95,29 +95,60 @@
-
+
预测结果
-
-
- {{ formatCurrency(predictionResult.predicted_cost) }}
-
-
- {{ formatCurrency(predictionResult.confidence_interval.lower) }} ~
- {{ formatCurrency(predictionResult.confidence_interval.upper) }}
-
-
+
+
+
+
机器学习模型预测
+
+
+ {{ getModelName(mlPrediction.model_info.type) }}
+
+
+ {{ mlPrediction.model_info.name }}
+
+
+ {{ formatMoney(mlPrediction.predicted_cost) }}
+
+
+ {{ formatMoney(mlPrediction.confidence_interval.lower) }} ~
+ {{ formatMoney(mlPrediction.confidence_interval.upper) }}
+
+
+
+
+
+
+
PLS回归预测
+
+
+ {{ getModelName(plsPrediction.model_info.type) }}
+
+
+ {{ plsPrediction.model_info.name }}
+
+
+ {{ formatMoney(plsPrediction.predicted_cost) }}
+
+
+ {{ formatMoney(plsPrediction.confidence_interval.lower) }} ~
+ {{ formatMoney(plsPrediction.confidence_interval.upper) }}
+
+
+
-
\ No newline at end of file
diff --git a/frontend/src/views/TrainingPage.vue b/frontend/src/views/TrainingPage.vue
index b4724bc..edd20c1 100644
--- a/frontend/src/views/TrainingPage.vue
+++ b/frontend/src/views/TrainingPage.vue
@@ -6,52 +6,49 @@
-
-
-
-
-
+
+
+
+
+
-
-
-
-
+
+
+ />
-
-
-
+
+ />
-
-
-
- XGBoost
- LightGBM
- GBDT
- Random Forest
+
+
+ PLS回归
+ XGBoost
+ LightGBM
+ GBM
+ Random Forest
-
-
- {{ training ? '训练中...' : '开始训练' }}
+
+ 开始训练
@@ -60,44 +57,56 @@
训练结果
-
-
-
+
+
+
最佳模型: {{ getModelName(trainingResult.best_model.type) }}
+
R²分数: {{ formatNumber(trainingResult.best_model.r2) }}
+
MAE: {{ formatNumber(trainingResult.best_model.mae) }} 元
+
RMSE: {{ formatNumber(trainingResult.best_model.rmse) }} 元
+
+
+
+
+
- {{ formatModelName(scope.row.model) }}
+ {{ getModelName(scope.row.model) }}
+
+
-
+
- {{ scope.row.train.r2.toFixed(4) }}
+ {{ formatNumber(scope.row.train.r2) }}
-
+
- {{ scope.row.train.mae.toFixed(2) }}
+ {{ formatNumber(scope.row.train.mae) }}
-
+
- {{ scope.row.train.rmse.toFixed(2) }}
+ {{ formatNumber(scope.row.train.rmse) }}
-
-
+
+
+
+
- {{ scope.row.validation.r2.toFixed(4) }}
+ {{ formatNumber(scope.row.validation.r2) }}
-
+
- {{ scope.row.validation.mae.toFixed(2) }}
+ {{ formatNumber(scope.row.validation.mae) }}
-
+
- {{ scope.row.validation.rmse.toFixed(2) }}
+ {{ formatNumber(scope.row.validation.rmse) }}
@@ -106,148 +115,142 @@
特征重要性
-
-
-
+
+
+
-
+ {{ formatNumber(scope.row.importance) }}
-
-
-
-
最佳模型
-
-
- {{ formatModelName(trainingResult.best_model.type) }}
-
-
- {{ trainingResult.best_model.r2.toFixed(4) }}
-
-
- {{ formatMoney(trainingResult.best_model.mae) }}
-
-
- {{ formatMoney(trainingResult.best_model.rmse) }}
-
-
-
@@ -274,21 +336,35 @@ onMounted(() => {
padding: 20px;
.training-card {
- max-width: 800px;
- margin: 0 auto;
- }
-
- .training-result {
- margin-top: 20px;
- padding: 20px;
- background-color: #f5f7fa;
- border-radius: 4px;
- }
-
- h3, h4 {
- margin: 20px 0;
- padding-left: 10px;
- border-left: 4px solid #409EFF;
+ .training-result {
+ margin-top: 20px;
+
+ .best-model-info {
+ background-color: #f5f7fa;
+ padding: 15px;
+ border-radius: 4px;
+ margin-bottom: 20px;
+ }
+
+ .feature-importance {
+ margin-top: 20px;
+
+ .importance-bar {
+ width: 100%;
+ background-color: #f5f7fa;
+ border-radius: 4px;
+
+ .importance-value {
+ background-color: #409eff;
+ color: white;
+ padding: 4px 8px;
+ border-radius: 4px;
+ text-align: right;
+ transition: width 0.3s ease;
+ }
+ }
+ }
+ }
}
}
\ No newline at end of file
diff --git a/run.py b/run.py
index d561cde..5bca7f6 100644
--- a/run.py
+++ b/run.py
@@ -1,61 +1,13 @@
-import os
+from src.app import create_app
import logging
-from src.app import app
-# 确保必要的目录存在
-def ensure_directories():
- """
- 确保所有必要的目录都存在
- """
- directories = [
- 'logs',
- 'data',
- 'models',
- 'uploads'
- ]
-
- for directory in directories:
- os.makedirs(directory, exist_ok=True)
+# 创建应用实例
+app = create_app()
-# 配置日志
-def setup_logging():
- """
- 配置日志系统
- """
- logging.basicConfig(
- filename='logs/server.log',
- level=logging.INFO,
- format='%(asctime)s - %(levelname)s - %(message)s'
- )
-
- # 同时输出到控制台
- console_handler = logging.StreamHandler()
- console_handler.setLevel(logging.INFO)
- formatter = logging.Formatter('%(asctime)s - %(levelname)s - %(message)s')
- console_handler.setFormatter(formatter)
- logging.getLogger('').addHandler(console_handler)
+if __name__ == '__main__':
+ # 设置日志
+ logging.basicConfig(level=logging.INFO)
+ logging.info('=== Server Starting ===')
+ logging.info('Initializing directories...')
-if __name__ == "__main__":
- try:
- # 初始化目录
- ensure_directories()
-
- # 设置日志
- setup_logging()
-
- # 记录启动信息
- logging.info("=== Server Starting ===")
- logging.info("Initializing directories...")
- logging.info("Setting up logging system...")
-
- # 启动服务器
- app.run(
- host='localhost',
- port=5001,
- debug=True,
- use_reloader=False # 禁用重载器以避免模型重复加载
- )
-
- except Exception as e:
- logging.error(f"Server failed to start: {str(e)}")
- raise
\ No newline at end of file
+ app.run(host='0.0.0.0', port=5001, debug=True)
\ No newline at end of file
diff --git a/src/api.py b/src/api.py
index e7a2d34..ee2fea3 100644
--- a/src/api.py
+++ b/src/api.py
@@ -1,5 +1,5 @@
from flask import Flask, request, jsonify
-from .model_training import ModelTrainer
+from .model_trainer import ModelTrainer
from .cost_prediction import CostPredictor
from .feature_analysis import FeatureAnalysis
import pandas as pd
diff --git a/src/app.py b/src/app.py
index c190efb..037fb79 100644
--- a/src/app.py
+++ b/src/app.py
@@ -1,68 +1,50 @@
from flask import Flask
from flask_cors import CORS
-import logging
-import os
from .routes import api_bp
+from .logger import setup_logger
+import os
+
+# 获取logger
+logger = setup_logger(__name__)
def create_app():
"""
创建并配置Flask应用
"""
- app = Flask(__name__)
-
- # 配置跨域
- CORS(app)
-
- # 配置日志
- setup_logging()
-
- # 注册蓝图
- app.register_blueprint(api_bp, url_prefix='/api')
-
- # 错误处理
- @app.errorhandler(404)
- def not_found_error(error):
- logging.error(f"404 error: {error}")
- return {'error': 'Resource not found'}, 404
+ try:
+ # 创建必要的目录
+ os.makedirs('logs', exist_ok=True)
+ os.makedirs('data', exist_ok=True)
+ os.makedirs('models', exist_ok=True)
- @app.errorhandler(500)
- def internal_error(error):
- logging.error(f"500 error: {error}")
- return {'error': 'Internal server error'}, 500
+ logger.info("=== Server Starting ===")
+ logger.info("Initializing directories...")
- @app.errorhandler(Exception)
- def handle_exception(error):
- logging.error(f"Unhandled exception: {error}", exc_info=True)
- return {'error': str(error)}, 500
-
- return app
+ # 创建Flask应用
+ app = Flask(__name__)
+
+ # 配置CORS
+ CORS(app)
+ logger.info("CORS enabled")
+
+ # 注册API蓝图
+ app.register_blueprint(api_bp, url_prefix='/api')
+ logger.info("API blueprint registered")
+
+ # 配置数据库连接
+ app.config['MYSQL_HOST'] = 'localhost'
+ app.config['MYSQL_USER'] = 'root'
+ app.config['MYSQL_PASSWORD'] = '123456'
+ app.config['MYSQL_DB'] = 'equipment_cost_db'
+
+ logger.info("Starting server...")
+
+ return app
+
+ except Exception as e:
+ logger.error(f"Error creating app: {str(e)}")
+ raise
-def setup_logging():
- """
- 配置日志系统
- """
- # 确保日志目录存在
- os.makedirs('logs', exist_ok=True)
-
- # 配置日志格式
- logging.basicConfig(
- filename='logs/api.log',
- level=logging.INFO,
- format='%(asctime)s - %(levelname)s - %(message)s'
- )
-
- # 同时输出到控制台
- console_handler = logging.StreamHandler()
- console_handler.setLevel(logging.INFO)
- formatter = logging.Formatter('%(asctime)s - %(levelname)s - %(message)s')
- console_handler.setFormatter(formatter)
- logging.getLogger('').addHandler(console_handler)
-
-app = create_app()
-
-@app.route('/health')
-def health_check():
- """
- 健康检查端点
- """
- return {'status': 'ok'}
\ No newline at end of file
+if __name__ == '__main__':
+ app = create_app()
+ app.run(host='localhost', port=5001)
\ No newline at end of file
diff --git a/src/cost_prediction.py b/src/cost_prediction.py
index 4305d86..66662a2 100644
--- a/src/cost_prediction.py
+++ b/src/cost_prediction.py
@@ -10,6 +10,9 @@ from .feature_analysis import FeatureAnalysis
import logging
from src.model_trainer import ModelTrainer
from src.database import get_db_connection
+from .logger import setup_logger
+
+logger = setup_logger(__name__)
class CostPredictor:
def __init__(self):
@@ -33,7 +36,7 @@ class CostPredictor:
def load_model(self):
"""
- 加载预训练模型和标准化器
+ 加载预训练型和标准化器
"""
try:
model_dir = 'models'
@@ -142,12 +145,13 @@ class CostPredictor:
def predict(self, data):
"""
- 预测成本
+ 使用训练好的最优模型进行预测
"""
try:
+ logger.info(f"Starting prediction for {data.get('type')}")
equipment_type = data.get('type')
- # 加载模型
+ # 加载已训练的最优模型
trainer = ModelTrainer()
if not trainer.load_model(equipment_type):
raise ValueError(f"No trained model found for {equipment_type}")
@@ -160,23 +164,22 @@ class CostPredictor:
y_pred = trainer.predict(X)
# 计算置信区间
- confidence_interval = self._calculate_confidence_interval(y_pred[0])
+ confidence_interval = trainer._calculate_confidence_interval(y_pred[0])
- # 确保预测值和置信区间都是正数且合理的范围
- predicted_cost = max(1000, float(y_pred[0])) # 最小值设为1000元
- lower_bound = max(1000, float(confidence_interval[0]))
- upper_bound = max(predicted_cost * 1.2, float(confidence_interval[1])) # 至少比预测值大20%
+ # 获取模型类型
+ model_type = trainer.get_model_type()
return {
- 'predicted_cost': predicted_cost,
+ 'predicted_cost': float(y_pred[0]),
+ 'model_type': model_type, # 返回使用的模型类型
'confidence_interval': {
- 'lower': lower_bound,
- 'upper': upper_bound
+ 'lower': float(confidence_interval[0]),
+ 'upper': float(confidence_interval[1])
}
}
except Exception as e:
- logging.error(f"Prediction error: {str(e)}")
+ logger.error(f"Prediction error: {str(e)}")
raise
def _calculate_confidence_interval(self, prediction, confidence=0.95):
@@ -215,4 +218,39 @@ class CostPredictor:
'mse': float(mean_squared_error(y_true, y_pred)),
'rmse': float(np.sqrt(mean_squared_error(y_true, y_pred))),
'r2': float(r2_score(y_true, y_pred))
- }
\ No newline at end of file
+ }
+
+ def predict_pls(self, data):
+ """
+ 使用 PLS 模型预测成本
+ """
+ try:
+ logger.info(f"Starting PLS prediction for {data.get('type')}")
+ equipment_type = data.get('type')
+
+ # 加载 PLS 模型
+ trainer = ModelTrainer()
+ if not trainer.load_model(equipment_type, model_type='pls'): # 指定加载 PLS 模型
+ raise ValueError(f"No trained PLS model found for {equipment_type}")
+
+ # 准备特征数据
+ features = self.feature_analyzer.get_equipment_specific_features(equipment_type)
+ X = np.array([[data.get(feature) for feature in features]])
+
+ # 预测
+ y_pred = trainer.predict(X)
+
+ # 计算置信区间
+ confidence_interval = trainer._calculate_confidence_interval(y_pred[0])
+
+ return {
+ 'predicted_cost': float(y_pred[0]),
+ 'confidence_interval': {
+ 'lower': float(confidence_interval[0]),
+ 'upper': float(confidence_interval[1])
+ }
+ }
+
+ except Exception as e:
+ logger.error(f"PLS prediction error: {str(e)}")
+ raise
\ No newline at end of file
diff --git a/src/create_template.py b/src/create_template.py
index 1b49aea..0f7d545 100644
--- a/src/create_template.py
+++ b/src/create_template.py
@@ -3,6 +3,9 @@ import openpyxl
from openpyxl.styles import PatternFill, Font, Alignment
from openpyxl.worksheet.datavalidation import DataValidation
import os
+from .logger import setup_logger
+
+logger = setup_logger(__name__)
def create_excel_template():
"""
diff --git a/src/data_preparation.py b/src/data_preparation.py
index ea9e47a..6380647 100644
--- a/src/data_preparation.py
+++ b/src/data_preparation.py
@@ -13,6 +13,9 @@ import json
import logging
from src.database.db_connection import get_db_connection
from sklearn.metrics import mean_absolute_error, mean_squared_error
+from .logger import setup_logger
+
+logger = setup_logger(__name__)
class DataPreparation:
def __init__(self):
@@ -25,13 +28,13 @@ class DataPreparation:
准备训练数据
"""
try:
- logging.info(f"Preparing training data for {equipment_type}")
- logging.info(f"Raw data size: {len(equipment_data)}")
+ logger.info(f"Preparing training data for {equipment_type}")
+ logger.info(f"Raw data size: {len(equipment_data)}")
# 如果输入已经是 numpy 数组,直接返回
if isinstance(equipment_data, np.ndarray):
X = equipment_data
- logging.info(f"Input is already numpy array with shape: {X.shape}")
+ logger.info(f"Input is already numpy array with shape: {X.shape}")
# 处理无效值
X = np.nan_to_num(X, nan=0.0, posinf=0.0, neginf=0.0)
@@ -65,9 +68,9 @@ class DataPreparation:
if cost > 0: # 只使用正数成本值
targets.append(cost)
else:
- logging.warning(f"Skipping non-positive cost value: {cost}")
+ logger.warning(f"Skipping non-positive cost value: {cost}")
except (ValueError, TypeError, KeyError):
- logging.error(f"Invalid cost value: {item.get('actual_cost')}")
+ logger.error(f"Invalid cost value: {item.get('actual_cost')}")
continue
# 转换为numpy数组
@@ -75,19 +78,25 @@ class DataPreparation:
y = np.array(targets, dtype=float)
# 记录原始数据范围
- logging.info(f"Original X range: min={X.min()}, max={X.max()}")
- logging.info(f"Original y range: min={y.min()}, max={y.max()}")
-
- # 处理无效值
- X = np.nan_to_num(X, nan=0.0, posinf=0.0, neginf=0.0)
+ logger.info(f"Raw X range: min={X.min()}, max={X.max()}")
+ logger.info(f"Raw y range: min={y.min()}, max={y.max()}")
# 标准化特征和目标值
X_scaled = self.feature_scaler.fit_transform(X)
y_scaled = self.target_scaler.fit_transform(y.reshape(-1, 1)).ravel()
# 记录标准化后的数据范围
- logging.info(f"Scaled X range: min={X_scaled.min()}, max={X_scaled.max()}")
- logging.info(f"Scaled y range: min={y_scaled.min()}, max={y_scaled.max()}")
+ logger.info(f"Scaled X range: min={X_scaled.min()}, max={X_scaled.max()}")
+ logger.info(f"Scaled y range: min={y_scaled.min()}, max={y_scaled.max()}")
+
+ # 记录标准化器参数
+ logger.info("Feature scaler params:")
+ logger.info(f"Mean: {self.feature_scaler.mean_}")
+ logger.info(f"Scale: {self.feature_scaler.scale_}")
+
+ logger.info("Target scaler params:")
+ logger.info(f"Mean: {self.target_scaler.mean_}")
+ logger.info(f"Scale: {self.target_scaler.scale_}")
return {
'X': X_scaled,
@@ -98,7 +107,7 @@ class DataPreparation:
}
except Exception as e:
- logging.error(f"Error in data preparation: {str(e)}")
+ logger.error(f"Error in data preparation: {str(e)}")
raise Exception(f"Training error: {str(e)}")
def prepare_validation_data(self, validation_data, equipment_type, feature_names=None, scalers=None):
@@ -106,13 +115,13 @@ class DataPreparation:
准备验证数据
"""
try:
- logging.info(f"Preparing validation data for {equipment_type}")
- logging.info(f"Raw validation data size: {len(validation_data)}")
+ logger.info(f"Preparing validation data for {equipment_type}")
+ logger.info(f"Raw validation data size: {len(validation_data)}")
# 如果输入已经是 numpy 数组,直接使用
if isinstance(validation_data, np.ndarray):
X = validation_data
- logging.info(f"Input is already numpy array with shape: {X.shape}")
+ logger.info(f"Input is already numpy array with shape: {X.shape}")
# 处理无效值
X = np.nan_to_num(X, nan=0.0, posinf=0.0, neginf=0.0)
@@ -123,9 +132,9 @@ class DataPreparation:
else:
X_scaled = X
- logging.info(f"Preprocessed data shape: {X_scaled.shape}")
- logging.info(f"Validation features shape: {X_scaled.shape}")
- logging.info(f"Validation features type: {X_scaled.dtype}")
+ logger.info(f"Preprocessed data shape: {X_scaled.shape}")
+ logger.info(f"Validation features shape: {X_scaled.shape}")
+ logger.info(f"Validation features type: {X_scaled.dtype}")
return {
'X': X_scaled,
@@ -153,13 +162,13 @@ class DataPreparation:
# 提取目标值(成本)并验证范围
try:
cost = float(item['actual_cost'])
- logging.info(f"Raw cost value: {cost}")
+ logger.info(f"Raw cost value: {cost}")
if cost > 0: # 只使用正数成本值
targets.append(cost)
else:
- logging.warning(f"Skipping non-positive cost value: {cost}")
+ logger.warning(f"Skipping non-positive cost value: {cost}")
except (ValueError, TypeError):
- logging.error(f"Invalid cost value: {item.get('actual_cost')}")
+ logger.error(f"Invalid cost value: {item.get('actual_cost')}")
continue
# 转换为numpy数组
@@ -167,8 +176,8 @@ class DataPreparation:
y = np.array(targets, dtype=float)
# 记录数据范围
- logging.info(f"Features range: min={X.min()}, max={X.max()}")
- logging.info(f"Targets range: min={y.min()}, max={y.max()}")
+ logger.info(f"Features range: min={X.min()}, max={X.max()}")
+ logger.info(f"Targets range: min={y.min()}, max={y.max()}")
# 处理无效值
X = np.nan_to_num(X, nan=0.0, posinf=0.0, neginf=0.0)
@@ -184,13 +193,23 @@ class DataPreparation:
X_scaled = X
y_scaled = y
- logging.info(f"Preprocessed data shape: {X_scaled.shape}")
- logging.info(f"Validation features shape: {X_scaled.shape}")
- logging.info(f"Validation features type: {X_scaled.dtype}")
+ logger.info(f"Preprocessed data shape: {X_scaled.shape}")
+ logger.info(f"Validation features shape: {X_scaled.shape}")
+ logger.info(f"Validation features type: {X_scaled.dtype}")
# 记录标准化后的数据范围
- logging.info(f"Scaled validation X range: min={X_scaled.min()}, max={X_scaled.max()}")
- logging.info(f"Scaled validation y range: min={y_scaled.min()}, max={y_scaled.max()}")
+ logger.info(f"Scaled validation X range: min={X_scaled.min()}, max={X_scaled.max()}")
+ logger.info(f"Scaled validation y range: min={y_scaled.min()}, max={y_scaled.max()}")
+
+ # 确保特征维度一致
+ if not feature_names:
+ feature_names = self.feature_analyzer.get_equipment_specific_features(equipment_type)
+
+ logger.info(f"Expected features: {len(feature_names)}")
+ logger.info(f"Actual features: {X_scaled.shape[1]}")
+
+ if X_scaled.shape[1] != len(feature_names):
+ raise ValueError(f"Feature dimension mismatch: expected {len(feature_names)}, got {X_scaled.shape[1]}")
return {
'X': X_scaled,
@@ -198,9 +217,9 @@ class DataPreparation:
}
except Exception as e:
- logging.error(f"Error in validation data preparation: {str(e)}")
- logging.error(f"Feature names: {feature_names}")
- logging.error(f"Equipment type: {equipment_type}")
+ logger.error(f"Error in validation data preparation: {str(e)}")
+ logger.error(f"Feature names: {feature_names}")
+ logger.error(f"Equipment type: {equipment_type}")
raise Exception(f"Validation error: {str(e)}")
def calculate_derived_features(self, data, equipment_type):
@@ -210,5 +229,5 @@ class DataPreparation:
try:
return self.feature_analyzer.calculate_derived_features(data, equipment_type)
except Exception as e:
- logging.error(f"Error calculating derived features: {str(e)}")
+ logger.error(f"Error calculating derived features: {str(e)}")
raise Exception(f"Feature calculation error: {str(e)}")
\ No newline at end of file
diff --git a/src/database/db_connection.py b/src/database/db_connection.py
index 2c3288e..361a4d2 100644
--- a/src/database/db_connection.py
+++ b/src/database/db_connection.py
@@ -1,28 +1,37 @@
import mysql.connector
from mysql.connector import Error
-import logging
from contextlib import contextmanager
+import os
+from dotenv import load_dotenv
+from ..logger import setup_logger
-# 数据库配置
-DB_CONFIG = {
- 'host': 'localhost',
- 'user': 'root',
- 'password': '123456',
- 'database': 'equipment_cost_db'
-}
+# 获取logger
+logger = setup_logger(__name__)
+
+# 加载环境变量
+load_dotenv()
@contextmanager
def get_db_connection():
"""
数据库连接上下文管理器
"""
- conn = None
+ connection = None
try:
- conn = mysql.connector.connect(**DB_CONFIG)
- yield conn
+ connection = mysql.connector.connect(
+ host=os.getenv('MYSQL_HOST', 'localhost'),
+ user=os.getenv('MYSQL_USER', 'root'),
+ password=os.getenv('MYSQL_PASSWORD', '123456'),
+ database=os.getenv('MYSQL_DATABASE', 'equipment_cost_db')
+ )
+ logger.info("Database connection established")
+ yield connection
+
except Error as e:
- logging.error(f"Error connecting to MySQL: {str(e)}")
+ logger.error(f"Error connecting to MySQL: {str(e)}")
raise
+
finally:
- if conn and conn.is_connected():
- conn.close()
\ No newline at end of file
+ if connection and connection.is_connected():
+ connection.close()
+ logger.info("Database connection closed")
\ No newline at end of file
diff --git a/src/feature_analysis.py b/src/feature_analysis.py
index de71a29..5b87972 100644
--- a/src/feature_analysis.py
+++ b/src/feature_analysis.py
@@ -5,6 +5,9 @@ from sklearn.preprocessing import StandardScaler
from sklearn.ensemble import RandomForestRegressor
from sklearn.metrics import r2_score
import logging
+from .logger import setup_logger
+
+logger = setup_logger(__name__)
class FeatureAnalysis:
def __init__(self):
@@ -182,7 +185,7 @@ class FeatureAnalysis:
return data
except Exception as e:
- logging.error(f"Error calculating derived features: {str(e)}")
+ logger.error(f"Error calculating derived features: {str(e)}")
raise
def analyze_features(self, features, target, feature_names):
@@ -235,7 +238,7 @@ class FeatureAnalysis:
}
except Exception as e:
- print(f"Error in feature analysis: {str(e)}")
+ logger.error(f"Error in feature analysis: {str(e)}")
raise
def preprocess_features(self, equipment_data, equipment_type):
@@ -258,9 +261,9 @@ class FeatureAnalysis:
mean_value = df[col].mean()
df[col] = df[col].fillna(mean_value)
- logging.info(f"Preprocessed data shape: {df.shape}")
+ logger.info(f"Preprocessed data shape: {df.shape}")
return df
except Exception as e:
- logging.error(f"Error preprocessing features: {str(e)}")
+ logger.error(f"Error preprocessing features: {str(e)}")
raise Exception(f"Feature preprocessing error: {str(e)}")
\ No newline at end of file
diff --git a/src/import_data.py b/src/import_data.py
index f4ff022..03dfa36 100644
--- a/src/import_data.py
+++ b/src/import_data.py
@@ -1,12 +1,8 @@
import pandas as pd
-import logging
+from .logger import setup_logger
from src.database.db_connection import get_db_connection
-logging.basicConfig(
- filename='logs/import.log',
- level=logging.INFO,
- format='%(asctime)s - %(levelname)s - %(message)s'
-)
+logger = setup_logger(__name__)
def import_training_data(excel_file):
"""
@@ -25,7 +21,7 @@ def import_training_data(excel_file):
cursor = conn.cursor()
# 1. 先导入火箭炮数据
- logging.info("开始导入火箭炮数据...")
+ logger.info("开始导入火箭炮数据...")
for _, row in rocket_df.iterrows():
equipment_names.add(row['名称'])
# 检查是否已存在相同名称的装备
@@ -36,7 +32,7 @@ def import_training_data(excel_file):
existing_equipment = cursor.fetchone()
if existing_equipment:
- logging.warning(f"火箭炮 '{row['名称']}' 已存在,跳过导入")
+ logger.warning(f"火箭炮 '{row['名称']}' 已存在,跳过导入")
continue
# 插入基本信息
@@ -96,15 +92,15 @@ def import_training_data(excel_file):
VALUES (%s, %s)
""", (equipment_id, row['成本_元']))
- logging.info("火箭炮数据导入完成")
+ logger.info("火箭炮数据导入完成")
# 2. 导入巡飞弹数据
- logging.info("开始导入巡飞弹数据...")
+ logger.info("开始导入巡飞弹数据...")
for index, row in missile_df.iterrows():
# 记录每行数据的空值情况
null_values = row[row.isna()].index.tolist()
if null_values:
- logging.info(f"行 {index + 2} 中的空值字段: {null_values}")
+ logger.info(f"行 {index + 2} 中的空值字段: {null_values}")
equipment_names.add(row['名称'])
# 检查是否已存在相同名称的装备
@@ -115,7 +111,7 @@ def import_training_data(excel_file):
existing_equipment = cursor.fetchone()
if existing_equipment:
- logging.warning(f"巡飞弹 '{row['名称']}' 已存在,跳过导入")
+ logger.warning(f"巡飞弹 '{row['名称']}' 已存在,跳过导入")
continue
# 插入基本信息
@@ -175,25 +171,25 @@ def import_training_data(excel_file):
VALUES (%s, %s)
""", (equipment_id, float(row['成本_元'])))
- logging.info("巡飞弹数据导入完成")
+ logger.info("巡飞弹数据导入完成")
# 提交之前的更改并关闭原有游标
cursor.close()
conn.commit()
# 3. 导入特殊参数
- logging.info("开始导入特殊参数...")
+ logger.info("开始导入特殊参数...")
for index, row in special_df.iterrows():
equipment_name = row['装备名称']
param_name = row['参数名称']
- logging.info(f"处理第 {index + 1} 条记录: 装备='{equipment_name}', 参数='{param_name}'")
+ logger.info(f"处理第 {index + 1} 条记录: 装备='{equipment_name}', 参数='{param_name}'")
if equipment_name not in equipment_names:
- logging.warning(f"未找到装备: {equipment_name},请检查名称是否正确")
+ logger.warning(f"未找到装备: {equipment_name},请检查名称是否正确")
continue
# 获取装备ID - 使用新的游标
- logging.debug(f"查询装备ID: {equipment_name}")
+ logger.debug(f"查询装备ID: {equipment_name}")
with conn.cursor() as id_cursor:
id_cursor.execute("""
SELECT id FROM equipment WHERE name = %s
@@ -201,14 +197,14 @@ def import_training_data(excel_file):
result = id_cursor.fetchone()
if not result:
- logging.warning(f"未找到装备: {equipment_name}")
+ logger.warning(f"未找到装备: {equipment_name}")
continue
equipment_id = result[0]
- logging.debug(f"找到装备ID: {equipment_id}")
+ logger.debug(f"找到装备ID: {equipment_id}")
# 检查参数是否存在 - 使用新的游标
- logging.debug(f"检查参数是否存在: equipment_id={equipment_id}, param_name='{param_name}'")
+ logger.debug(f"检查参数是否存在: equipment_id={equipment_id}, param_name='{param_name}'")
with conn.cursor() as check_cursor:
check_cursor.execute("""
SELECT id FROM custom_params
@@ -217,7 +213,7 @@ def import_training_data(excel_file):
exists = check_cursor.fetchone()
if exists:
- logging.warning(f"装备 '{equipment_name}' 的参数 '{param_name}' 已存在,跳过导入")
+ logger.warning(f"装备 '{equipment_name}' 的参数 '{param_name}' 已存在,跳过导入")
continue
# 插入新的参数 - 使用新的游标
@@ -225,7 +221,7 @@ def import_training_data(excel_file):
param_unit = row['参数单位'] if pd.notna(row['参数单位']) else None
param_desc = row['参数说明'] if pd.notna(row['参数说明']) else None
- logging.debug(f"插入新参数: value='{param_value}', unit='{param_unit}', desc='{param_desc}'")
+ logger.debug(f"插入新参数: value='{param_value}', unit='{param_unit}', desc='{param_desc}'")
with conn.cursor() as insert_cursor:
insert_cursor.execute("""
INSERT INTO custom_params
@@ -238,22 +234,22 @@ def import_training_data(excel_file):
param_unit,
param_desc
))
- logging.debug(f"成功插入参数记录")
+ logger.debug(f"成功插入参数记录")
# 最终提交
conn.commit()
- logging.info("特殊参数导入完成")
- logging.info("所有数据导入成功")
+ logger.info("特殊参数导入完成")
+ logger.info("所有数据导入成功")
return True
except Exception as e:
- logging.error(f"Error importing data: {str(e)}")
+ logger.error(f"Error importing data: {str(e)}")
raise
if __name__ == "__main__":
try:
excel_file = 'data/equipment_data_20241108.xlsx'
import_training_data(excel_file)
- logging.info("All data imported successfully")
+ logger.info("All data imported successfully")
except Exception as e:
- logging.error(f"Import failed: {str(e)}")
\ No newline at end of file
+ logger.error(f"Import failed: {str(e)}")
\ No newline at end of file
diff --git a/src/logger.py b/src/logger.py
new file mode 100644
index 0000000..52bb879
--- /dev/null
+++ b/src/logger.py
@@ -0,0 +1,33 @@
+import logging
+import os
+from datetime import datetime
+
+def setup_logger(name):
+ """
+ 创建并配置logger
+ """
+ # 创建logger
+ logger = logging.getLogger(name)
+
+ # 如果logger已经有处理器,直接返回
+ if logger.handlers:
+ return logger
+
+ # 设置日志级别
+ logger.setLevel(logging.INFO)
+
+ # 确保日志目录存在
+ os.makedirs('logs', exist_ok=True)
+
+ # 创建文件处理器
+ file_handler = logging.FileHandler('logs/api.log')
+ file_handler.setLevel(logging.INFO)
+
+ # 创建格式化器
+ formatter = logging.Formatter('%(asctime)s - %(name)s - %(levelname)s - %(message)s')
+ file_handler.setFormatter(formatter)
+
+ # 添加处理器
+ logger.addHandler(file_handler)
+
+ return logger
\ No newline at end of file
diff --git a/src/model_trainer.py b/src/model_trainer.py
index 94a5d70..660c38e 100644
--- a/src/model_trainer.py
+++ b/src/model_trainer.py
@@ -14,39 +14,64 @@ from datetime import datetime
import json
from src.database import get_db_connection
from src.data_preparation import DataPreparation
+from sklearn.cross_decomposition import PLSRegression
+from .logger import setup_logger
+
+logger = setup_logger(__name__)
class ModelTrainer:
def __init__(self):
+ """
+ 初始化 ModelTrainer
+ """
self.models = {
'xgboost': self._create_xgboost_model(),
'lightgbm': self._create_lightgbm_model(),
- 'gbdt': self._create_gbdt_model(),
- 'rf': self._create_rf_model()
+ 'gbm': self._create_gbm_model(),
+ 'rf': self._create_rf_model(),
+ 'pls': self._create_pls_model()
}
self.best_model = None
self.imputer = SimpleImputer(strategy='mean')
self.feature_scaler = None
self.target_scaler = None
+ self.equipment_type = None
+ self.feature_analyzer = FeatureAnalysis()
def fit_model(self, X_train, y_train, model_names, X_val=None, y_val=None, equipment_type=None):
"""
训练模型并返回评估结果
"""
try:
- # 记录数据范围
- logging.info(f"Training data range - X: min={X_train.min()}, max={X_train.max()}")
- logging.info(f"Training data range - y: min={y_train.min()}, max={y_train.max()}")
+ self.equipment_type = equipment_type
+ logger.info(f"Training data range - X: min={X_train.min()}, max={X_train.max()}")
+ logger.info(f"Training data range - y: min={y_train.min()}, max={y_train.max()}")
results = {}
best_score = -float('inf')
best_model_info = None
+ # 首先训练 PLS 模型
+ logger.info("Training pls...")
+ pls_model = self.models['pls']
+ pls_model.fit(X_train, y_train)
+ pls_metrics = self._calculate_metrics(
+ pls_model,
+ X_train, y_train,
+ X_val, y_val
+ )
+ results['pls'] = pls_metrics
+
+ # 训练其他机器学习模型
for model_name in model_names:
- if model_name not in self.models:
- logging.warning(f"Unknown model: {model_name}")
+ if model_name == 'pls': # 跳过 PLS 模型,因为已经训练过了
continue
- logging.info(f"Training {model_name}...")
+ if model_name not in self.models:
+ logger.warning(f"Unknown model: {model_name}")
+ continue
+
+ logger.info(f"Training {model_name}...")
model = self.models[model_name]
# 训练模型
@@ -59,48 +84,30 @@ class ModelTrainer:
X_val, y_val
)
- # 更新最佳模型
+ results[model_name] = metrics
+
+ # 更新最佳模型(只在机器学习模型中比较)
if metrics['validation']['r2'] > best_score:
best_score = metrics['validation']['r2']
- self.best_model = model
best_model_info = {
'type': model_name,
- 'r2': float(metrics['validation']['r2']),
- 'mae': float(metrics['validation']['mae']) if metrics['validation']['mae'] is not None else None,
- 'rmse': float(metrics['validation']['rmse']) if metrics['validation']['rmse'] is not None else None
+ 'r2': metrics['validation']['r2'],
+ 'mae': metrics['validation']['mae'],
+ 'rmse': metrics['validation']['rmse']
}
-
- # 转换 numpy 数据类型为 Python 原生类型
- results[model_name] = {
- 'train': {
- 'r2': float(metrics['train']['r2']),
- 'mae': float(metrics['train']['mae']),
- 'rmse': float(metrics['train']['rmse'])
- },
- 'validation': {
- 'r2': float(metrics['validation']['r2']),
- 'mae': float(metrics['validation']['mae']) if metrics['validation']['mae'] is not None else None,
- 'rmse': float(metrics['validation']['rmse']) if metrics['validation']['rmse'] is not None else None
- }
- }
+ self.best_model = model
- # 保存最佳模型
+ # 保存最佳模型和 PLS 模型
if equipment_type and best_model_info:
- self._save_best_model(equipment_type, best_model_info, X_train)
-
- # 转换特征重要性为列表
- feature_importance = None
- if self.best_model and hasattr(self.best_model, 'feature_importances_'):
- feature_importance = [float(x) for x in self.best_model.feature_importances_]
+ self._save_best_model(equipment_type, best_model_info, X_train, y_train, X_val, y_val)
return {
'metrics': results,
- 'best_model': best_model_info,
- 'feature_importance': feature_importance
+ 'best_model': best_model_info
}
except Exception as e:
- logging.error(f"Error in model training: {str(e)}")
+ logger.error(f"Error in model training: {str(e)}")
raise
def _calculate_metrics(self, model, X_train, y_train, X_val=None, y_val=None):
@@ -120,8 +127,8 @@ class ModelTrainer:
y_train_orig = y_train
# 记录预测范围
- logging.info(f"Train predictions range: min={train_pred.min()}, max={train_pred.max()}")
- logging.info(f"Train actual range: min={y_train_orig.min()}, max={y_train_orig.max()}")
+ logger.info(f"Train predictions range: min={train_pred.min()}, max={train_pred.max()}")
+ logger.info(f"Train actual range: min={y_train_orig.min()}, max={y_train_orig.max()}")
train_metrics = {
'r2': r2_score(y_train_orig, train_pred),
@@ -141,8 +148,8 @@ class ModelTrainer:
y_val_orig = y_val
# 记录预测范围
- logging.info(f"Validation predictions range: min={val_pred.min()}, max={val_pred.max()}")
- logging.info(f"Validation actual range: min={y_val_orig.min()}, max={y_val_orig.max()}")
+ logger.info(f"Validation predictions range: min={val_pred.min()}, max={val_pred.max()}")
+ logger.info(f"Validation actual range: min={y_val_orig.min()}, max={y_val_orig.max()}")
val_metrics = {
'r2': r2_score(y_val_orig, val_pred),
@@ -169,9 +176,9 @@ class ModelTrainer:
"""
return xgb.XGBRegressor(
n_estimators=50, # 减少树的数量
- learning_rate=0.05, # 减小学习率
- max_depth=3, # 减小树的深度
- min_child_weight=3, # 增加最小子节点权重
+ learning_rate=0.05, # 学习率
+ max_depth=3, # 减小树的深
+ min_child_weight=3, # 增加节点权重
subsample=0.7, # 减小样本采样比例
colsample_bytree=0.7, # 减小特征采样比例
reg_alpha=0.1, # L1 正则化
@@ -198,19 +205,18 @@ class ModelTrainer:
verbose=-1
)
- def _create_gbdt_model(self):
+ def _create_gbm_model(self):
"""
- 创建 GBDT 模型,增强正则化以减轻过拟合
+ 创建 GBM 模型,增强正则化以减轻过拟合
"""
return GradientBoostingRegressor(
- n_estimators=20, # 减少树的数量
- learning_rate=0.01, # 减小学习率
- max_depth=2, # 减小树的深度
- min_samples_split=4, # 增加分裂所需的最小样本数
- min_samples_leaf=3, # 增加叶子节点最小样本数
- subsample=0.5, # 减小样本采样比例
+ n_estimators=100,
+ learning_rate=0.1,
+ max_depth=3,
random_state=42,
- validation_fraction=0.2 # 使用部分训练数据作为验证集
+ subsample=0.8,
+ min_samples_split=3,
+ min_samples_leaf=2
)
def _create_rf_model(self):
@@ -218,26 +224,34 @@ class ModelTrainer:
创建随机森林模型,针对小样本数据调整参数
"""
return RandomForestRegressor(
- n_estimators=100, # 增加树的数量
- max_depth=4, # 限制树的深度
- min_samples_split=2, # 减小分需的最小样本数
- min_samples_leaf=1, # 减小叶子节点最小样本数
- max_features='sqrt', # 特征采样
- bootstrap=True, # 使用 bootstrap 采样
- oob_score=True, # 计算袋外分数
- random_state=42
+ n_estimators=100,
+ max_depth=3,
+ random_state=42,
+ min_samples_split=3,
+ min_samples_leaf=2
)
- def _save_best_model(self, equipment_type, best_model_info, X_train):
+ def _create_pls_model(self):
"""
- 保存最佳模型
+ 创建 PLS 模型,优化参数配置
+ """
+ return PLSRegression(
+ n_components=2, # 减少主成分数量,从5减到2
+ scale=True, # 保持数据标准化
+ max_iter=500, # 减少最大迭代次数,避免过拟合
+ tol=1e-6 # 降低收敛精度,避免过拟合
+ )
+
+ def _save_best_model(self, equipment_type, best_model_info, X_train, y_train, X_val=None, y_val=None):
+ """
+ 保存最佳模型和 PLS 模型
"""
try:
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
model_dir = 'models'
os.makedirs(model_dir, exist_ok=True)
- # 保存模型文件
+ # 1. 保存最佳机器学习模型
model_path = f'{model_dir}/{equipment_type}_{timestamp}'
if isinstance(self.best_model, xgb.XGBRegressor):
self.best_model.save_model(f'{model_path}.json')
@@ -246,128 +260,180 @@ class ModelTrainer:
joblib.dump(self.best_model, f'{model_path}.joblib')
model_format = 'joblib'
- # 验证标准化器
- if not isinstance(self.feature_scaler, StandardScaler):
- raise ValueError("Invalid feature scaler")
- if not isinstance(self.target_scaler, StandardScaler):
- raise ValueError("Invalid target scaler")
-
- # 保存标准化器
+ # 2. 保存 PLS 模型
+ pls_model = self.models['pls']
+ pls_path = f'{model_dir}/{equipment_type}_{timestamp}_pls.joblib'
+ joblib.dump(pls_model, pls_path)
+
+ # 3. 保存标准化器
scaler_path = f'{model_dir}/{equipment_type}_{timestamp}_scaler.joblib'
joblib.dump({
'feature_scaler': self.feature_scaler,
'target_scaler': self.target_scaler
}, scaler_path)
- logging.info(f"Saved model to {model_path}.{model_format}")
- logging.info(f"Saved scalers to {scaler_path}")
+ logger.info(f"Saved best model to {model_path}.{model_format}")
+ logger.info(f"Saved PLS model to {pls_path}")
+ logger.info(f"Saved scalers to {scaler_path}")
- # 更新数据库中的模型记录
+ # 4. 更新数据库中的模型记录
with get_db_connection() as conn:
cursor = conn.cursor()
- # 将之前的激活模型设置为非激活
+ # 将所有模型设置为非激活
cursor.execute("""
UPDATE trained_models
SET is_active = FALSE
WHERE equipment_type = %s
""", (equipment_type,))
- # 插入新的模型记录
+ # 获取 PLS 模型的评估指标
+ pls_metrics = self._calculate_metrics(
+ self.models['pls'],
+ X_train,
+ y_train,
+ X_val,
+ y_val
+ )
+
+ # 保存最佳机器学习模型记录
+ self.best_model.equipment_type = equipment_type # 设置装备类型
+ ml_feature_importance = self._get_feature_importance(self.best_model)
+
cursor.execute("""
INSERT INTO trained_models (
- model_name, model_type, equipment_type, model_path,
- scaler_path, r2_score, mae, rmse, feature_importance,
- training_data_size, training_date, is_active, created_by
- ) VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, NOW(), TRUE, 'system')
+ model_name, model_type, equipment_type, model_path, scaler_path,
+ r2_score, mae, rmse, feature_importance, training_data_size,
+ training_date, is_active, created_by
+ ) VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, NOW(), TRUE, %s)
""", (
- f'{best_model_info["type"]}_{timestamp}',
- best_model_info["type"],
- equipment_type,
- f'{model_path}.{model_format}',
- scaler_path,
- best_model_info["r2"],
- best_model_info["mae"],
- best_model_info["rmse"],
- json.dumps(self.feature_importance) if hasattr(self, 'feature_importance') else None,
- len(X_train)
+ f"{equipment_type}_{timestamp}", # model_name
+ best_model_info['type'], # model_type
+ equipment_type, # equipment_type
+ f"{model_path}.{model_format}", # model_path
+ scaler_path, # scaler_path
+ best_model_info['r2'], # r2_score
+ best_model_info['mae'], # mae
+ best_model_info['rmse'], # rmse
+ json.dumps(ml_feature_importance), # feature_importance
+ len(X_train), # training_data_size
+ 'system' # created_by
+ ))
+
+ # 保存 PLS 模型记录
+ pls_model.equipment_type = equipment_type # 设置装备类型
+ pls_feature_importance = self._get_feature_importance(pls_model)
+
+ cursor.execute("""
+ INSERT INTO trained_models (
+ model_name, model_type, equipment_type, model_path, scaler_path,
+ r2_score, mae, rmse, feature_importance, training_data_size,
+ training_date, is_active, created_by
+ ) VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, NOW(), TRUE, %s)
+ """, (
+ f"{equipment_type}_{timestamp}_pls", # model_name
+ 'pls', # model_type
+ equipment_type, # equipment_type
+ pls_path, # model_path
+ scaler_path, # scaler_path
+ float(pls_metrics['validation']['r2']), # r2_score
+ float(pls_metrics['validation']['mae']), # mae
+ float(pls_metrics['validation']['rmse']), # rmse
+ json.dumps(pls_feature_importance), # feature_importance
+ len(X_train), # training_data_size
+ 'system' # created_by
))
conn.commit()
- logging.info(f"Best model saved: {model_path}")
- return True
-
except Exception as e:
- logging.error(f"Error saving best model: {str(e)}")
- return False
+ logger.error(f"Error saving models: {str(e)}")
+ logger.error("Detailed traceback:", exc_info=True)
+ raise
- def load_model(self, equipment_type):
+ def load_model(self, equipment_type, model_type='ml'):
"""
加载已训练的模型
"""
try:
- logging.info(f"Loading model for {equipment_type}")
+ logger.info(f"Loading {model_type} model for {equipment_type}")
- # 从数据库获最新的激活模型
+ # 从数据库获取激活的模型
with get_db_connection() as conn:
cursor = conn.cursor(dictionary=True)
- cursor.execute("""
- SELECT * FROM trained_models
- WHERE equipment_type = %s AND is_active = TRUE
- ORDER BY training_date DESC LIMIT 1
- """, (equipment_type,))
+ # 构建查询语句
+ if model_type == 'pls':
+ query = """
+ SELECT * FROM trained_models
+ WHERE equipment_type = %s
+ AND model_type = 'pls'
+ AND is_active = TRUE
+ LIMIT 1
+ """
+ params = (equipment_type,)
+ else:
+ query = """
+ SELECT * FROM trained_models
+ WHERE equipment_type = %s
+ AND model_type != 'pls'
+ AND is_active = TRUE
+ LIMIT 1
+ """
+ params = (equipment_type,)
+
+ # 记录查询信息
+ logger.info(f"Executing query: {query}")
+ logger.info(f"Query params: {params}")
+
+ cursor.execute(query, params)
model_record = cursor.fetchone()
- if not model_record:
- raise ValueError(f"No active model found for {equipment_type}")
- logging.info(f"Found model: {model_record['model_name']}")
- logging.info(f"Model path: {model_record['model_path']}")
- logging.info(f"Scaler path: {model_record['scaler_path']}")
+ # 记录查询结果
+ if model_record:
+ logger.info(f"Found model record: {model_record}")
+ else:
+ logger.warning(f"No active model found for type {model_type}")
+ return False
# 检查文件是否存在
+ logger.info(f"Checking model file: {model_record['model_path']}")
+ logger.info(f"Checking scaler file: {model_record['scaler_path']}")
+
if not os.path.exists(model_record['model_path']):
+ logger.error(f"Model file not found: {model_record['model_path']}")
raise FileNotFoundError(f"Model file not found: {model_record['model_path']}")
+
if not os.path.exists(model_record['scaler_path']):
+ logger.error(f"Scaler file not found: {model_record['scaler_path']}")
raise FileNotFoundError(f"Scaler file not found: {model_record['scaler_path']}")
# 加载模型文件
- if model_record['model_type'] == 'xgboost':
- self.best_model = xgb.XGBRegressor()
- self.best_model.load_model(model_record['model_path'])
- else:
+ logger.info(f"Loading model from {model_record['model_path']}")
+ if model_type == 'pls':
self.best_model = joblib.load(model_record['model_path'])
+ logger.info("Loaded PLS model")
+ else:
+ if model_record['model_type'] == 'xgboost':
+ self.best_model = xgb.XGBRegressor()
+ self.best_model.load_model(model_record['model_path'])
+ logger.info("Loaded XGBoost model")
+ else:
+ self.best_model = joblib.load(model_record['model_path'])
+ logger.info(f"Loaded {model_record['model_type']} model")
# 加载标准化器
- try:
- scalers = joblib.load(model_record['scaler_path'])
- logging.info(f"Loaded scalers: {scalers.keys()}")
-
- if 'feature_scaler' not in scalers or 'target_scaler' not in scalers:
- raise ValueError("Missing scalers in saved file")
-
- self.feature_scaler = scalers['feature_scaler']
- self.target_scaler = scalers['target_scaler']
-
- # 验证标准化器
- if not hasattr(self.feature_scaler, 'transform') or not hasattr(self.target_scaler, 'transform'):
- raise ValueError("Invalid scaler objects")
-
- logging.info("Model and scalers loaded successfully")
- logging.info(f"Feature scaler type: {type(self.feature_scaler)}")
- logging.info(f"Target scaler type: {type(self.target_scaler)}")
-
- except Exception as e:
- logging.error(f"Error loading scalers: {str(e)}")
- logging.error(f"Scaler file content: {scalers if 'scalers' in locals() else 'Not loaded'}")
- raise ValueError(f"Failed to load scalers: {str(e)}")
+ logger.info(f"Loading scalers from {model_record['scaler_path']}")
+ scalers = joblib.load(model_record['scaler_path'])
+ self.feature_scaler = scalers['feature_scaler']
+ self.target_scaler = scalers['target_scaler']
+ logger.info("Loaded scalers successfully")
return True
except Exception as e:
- logging.error(f"Error loading model: {str(e)}")
- logging.error("Detailed traceback:", exc_info=True)
+ logger.error(f"Error loading model: {str(e)}")
+ logger.error(f"Detailed traceback:", exc_info=True)
return False
def predict(self, features):
@@ -384,39 +450,163 @@ class ModelTrainer:
if not self.target_scaler:
raise ValueError("Target scaler not loaded")
- logging.info("Starting prediction")
- logging.info(f"Input features shape: {features.shape}")
- logging.info(f"Input features: \n{features}")
+ logger.info("Starting prediction")
+ logger.info(f"Input features shape: {features.shape}")
+ logger.info(f"Input features: \n{features}")
# 处理缺失值
features_filled = np.array(features, dtype=float)
features_filled[np.isnan(features_filled)] = 0
features_filled = np.nan_to_num(features_filled, 0)
- logging.info(f"Filled features: \n{features_filled}")
+ logger.info(f"Filled features: \n{features_filled}")
# 标准化特征
X = self.feature_scaler.transform(features_filled)
- logging.info(f"Transformed features shape: {X.shape}")
- logging.info(f"Transformed features: \n{X}")
+ logger.info(f"Transformed features shape: {X.shape}")
+ logger.info(f"Transformed features: \n{X}")
# 预测
y_pred_scaled = self.best_model.predict(X)
- logging.info(f"Scaled prediction shape: {y_pred_scaled.shape}")
- logging.info(f"Scaled prediction: {y_pred_scaled}")
+ logger.info(f"Scaled prediction shape: {y_pred_scaled.shape}")
+ logger.info(f"Scaled prediction: {y_pred_scaled}")
- # 反标准化
+ # ��标准化
y_pred = self.target_scaler.inverse_transform(y_pred_scaled.reshape(-1, 1))
- logging.info(f"Final prediction shape: {y_pred.shape}")
- logging.info(f"Final prediction: {y_pred}")
+ logger.info(f"Final prediction shape: {y_pred.shape}")
+ logger.info(f"Final prediction: {y_pred}")
# 记录标准化器的参数
- logging.info("Target scaler params:")
- logging.info(f"Mean: {self.target_scaler.mean_}")
- logging.info(f"Scale: {self.target_scaler.scale_}")
+ logger.info("Target scaler params:")
+ logger.info(f"Mean: {self.target_scaler.mean_}")
+ logger.info(f"Scale: {self.target_scaler.scale_}")
return y_pred.ravel()
except Exception as e:
- logging.error(f"Error in prediction: {str(e)}")
- raise
\ No newline at end of file
+ logger.error(f"Error in prediction: {str(e)}")
+ raise
+
+ def _get_feature_importance(self, model):
+ """
+ 获取特征重要性
+ """
+ try:
+ if not model:
+ return {}
+
+ # 获取特征名称
+ feature_analyzer = FeatureAnalysis()
+ feature_names = feature_analyzer.get_equipment_specific_features(self.equipment_type)
+
+ # 获取特���重要性
+ if hasattr(model, 'feature_importances_'):
+ importances = model.feature_importances_
+ elif hasattr(model, 'coef_'):
+ if len(model.coef_.shape) > 1: # 如果是二维数组
+ importances = np.abs(model.coef_[0]) # 取第一行
+ else:
+ importances = np.abs(model.coef_)
+ else:
+ return {}
+
+ # 创建特征重要性字典
+ importance_dict = {}
+ for name, importance in zip(feature_names, importances):
+ importance_dict[name] = float(importance) # 确保转换为 Python 标量
+
+ # 按重要性降序排序
+ sorted_dict = dict(sorted(
+ importance_dict.items(),
+ key=lambda x: x[1],
+ reverse=True
+ ))
+
+ # 过滤掉重要性为0的特征
+ return {k: v for k, v in sorted_dict.items() if v > 0}
+
+ except Exception as e:
+ logger.error(f"Error getting feature importance: {str(e)}")
+ return {}
+
+ def _calculate_confidence_interval(self, prediction, confidence=0.95):
+ """
+ 计算预测值的置信区间
+ """
+ try:
+ # 使用预测值的20%作为标准差(增加不确定性)
+ std = abs(prediction) * 0.2
+
+ # 计算置信区间
+ from scipy import stats
+ interval = stats.norm.interval(confidence, loc=prediction, scale=std)
+
+ # 确保区间值为正数且合理
+ lower = max(1000, interval[0]) # 最小值设为1000元
+ upper = max(prediction * 1.2, interval[1]) # 至少比预测值大20%
+
+ logger.info(f"Calculated confidence interval: [{lower:.2f}, {upper:.2f}]")
+
+ return [lower, upper]
+
+ except Exception as e:
+ logger.error(f"Error calculating confidence interval: {str(e)}")
+ # 如果计算失败,返回基于20%的简单区间
+ lower = max(1000, prediction * 0.8)
+ upper = prediction * 1.2
+ return [lower, upper]
+
+ def get_model_type(self):
+ """
+ 获取当前模型的类型
+ """
+ if isinstance(self.best_model, xgb.XGBRegressor):
+ return 'xgboost'
+ elif isinstance(self.best_model, lgb.LGBMRegressor):
+ return 'lightgbm'
+ elif isinstance(self.best_model, GradientBoostingRegressor):
+ return 'gbm'
+ elif isinstance(self.best_model, RandomForestRegressor):
+ return 'rf'
+ else:
+ return 'unknown'
+
+ def _get_pls_feature_importance(self):
+ """
+ 获取 PLS 模型的特征重要性
+ """
+ try:
+ if not self.models['pls']:
+ return {}
+
+ # 获取特征名称
+ feature_analyzer = FeatureAnalysis()
+ feature_names = feature_analyzer.get_equipment_specific_features(self.equipment_type)
+
+ # 获取 PLS 模型的系数作为特征重要性
+ pls_model = self.models['pls']
+ if hasattr(pls_model, 'coef_'):
+ # 使用绝对值作为重要性指标
+ importances = np.abs(pls_model.coef_.ravel()) # 使用 ravel() 展平数组
+ else:
+ return {}
+
+ # 创建特征重要性字典
+ importance_dict = {}
+ for name, importance in zip(feature_names, importances):
+ importance_dict[name] = float(importance) # 确保转换为 Python 标量
+
+ # 按重要性降序排序
+ sorted_dict = dict(sorted(
+ importance_dict.items(),
+ key=lambda x: x[1],
+ reverse=True
+ ))
+
+ # 过滤掉重要性为0的特征
+ return {k: v for k, v in sorted_dict.items() if v > 0}
+
+ except Exception as e:
+ logger.error(f"Error getting PLS feature importance: {str(e)}")
+ logger.error("Detailed traceback:", exc_info=True)
+ return {}
\ No newline at end of file
diff --git a/src/pls_regression.py b/src/pls_regression.py
deleted file mode 100644
index 25fa2c0..0000000
--- a/src/pls_regression.py
+++ /dev/null
@@ -1,313 +0,0 @@
-# -*- coding: utf-8 -*-
-from sklearn.cross_decomposition import PLSRegression
-from sklearn.preprocessing import StandardScaler
-import numpy as np
-import pandas as pd
-import logging
-from sklearn.metrics import r2_score, mean_absolute_error
-from sklearn.model_selection import LeaveOneOut
-import os
-from datetime import datetime
-import joblib
-from src.database.db_connection import get_db_connection
-
-class PLSPredictor:
- def __init__(self, n_components=2):
- """
- 初始化PLS回归模型
- """
- self.model = PLSRegression(
- n_components=n_components,
- scale=True,
- max_iter=500,
- tol=1e-6
- )
- self.scaler_X = StandardScaler()
- self.scaler_y = StandardScaler()
- self.feature_names = None
- self.model_path = None
-
- # 尝试加载已训练的模型
- self.load_model()
-
- # 初始化示例数据
- self._initialize_scalers()
-
- def _initialize_scalers(self):
- """
- 使用示例数据初始化标准化器
- """
- # 创建示例数据
- example_data = pd.DataFrame({
- 'length_m': [0.56, 0.58, 0.54],
- 'width_m': [0.15, 0.16, 0.14],
- 'height_m': [0.20, 0.21, 0.19],
- 'weight_kg': [2.72, 2.85, 2.60],
- 'max_range_km': [24, 26, 22],
- 'max_speed_kmh': [160.93, 170, 155],
- 'cruise_speed_kmh': [96.56, 100, 93],
- 'flight_time_min': [15, 16, 14],
- 'folded_length_mm': [560, 580, 540],
- 'folded_width_mm': [150, 160, 140],
- 'folded_height_mm': [200, 210, 190]
- })
-
- # 初始化特征标准化器
- self.scaler_X.fit(example_data)
-
- # 初始化目标变量标准化器
- example_costs = np.array([[1000000], [1100000], [900000]])
- self.scaler_y.fit(example_costs)
-
- def predict(self, features):
- """
- 使用PLS模型进行预测
- """
- try:
- # 转换输入数据为DataFrame
- if not isinstance(features, pd.DataFrame):
- features = pd.DataFrame([features])
-
- # 选择数值特征
- numeric_features = features.select_dtypes(include=[np.number]).columns
- X = features[numeric_features]
-
- # 标准化特征
- X_scaled = self.scaler_X.transform(X)
-
- # 预测
- y_pred_scaled = self.model.predict(X_scaled)
- y_pred = self.scaler_y.inverse_transform(y_pred_scaled)
-
- # 计算置信区间
- ci = self._calculate_confidence_intervals(y_pred)
-
- return {
- 'predicted_cost': float(abs(y_pred[0][0])),
- 'confidence_interval': {
- 'lower': float(abs(ci['lower'])),
- 'upper': float(abs(ci['upper']))
- }
- }
-
- except Exception as e:
- logging.error(f"Error in PLS prediction: {str(e)}")
- raise Exception(f"PLS prediction error: {str(e)}")
-
- def fit(self, X, y):
- """
- 训练PLS模型
- """
- try:
- logging.info("=== PLS Training Start ===")
-
- # 1. 检查输入数据
- logging.info(f"Input X type: {type(X)}, shape: {X.shape if hasattr(X, 'shape') else 'no shape'}")
- logging.info(f"Input y type: {type(y)}, shape: {y.shape if hasattr(y, 'shape') else 'no shape'}")
- logging.info(f"X data:\n{X}")
- logging.info(f"y data:\n{y}")
-
- # 2. 转换为numpy数组
- if isinstance(X, pd.DataFrame):
- # 保存特征名称
- self.feature_names = X.columns.tolist()
- X = X.values
- X = np.array(X, dtype=float)
- y = np.array(y, dtype=float)
-
- # 3. 标准化数据
- logging.info("Standardizing data...")
- X_scaled = self.scaler_X.fit_transform(X)
- y_scaled = self.scaler_y.fit_transform(y.reshape(-1, 1))
- logging.info(f"X_scaled shape: {X_scaled.shape}")
- logging.info(f"y_scaled shape: {y_scaled.shape}")
-
- # 4. 训练模型
- logging.info("Training PLS model...")
- self.model.fit(X_scaled, y_scaled.ravel())
- logging.info("PLS model training completed")
-
- # 5. 计算R²分数
- logging.info("Calculating R² score...")
- y_pred = self.model.predict(X_scaled)
- y_pred = self.scaler_y.inverse_transform(y_pred.reshape(-1, 1))
- r2 = r2_score(y.reshape(-1, 1), y_pred)
- logging.info(f"R² score: {r2}")
-
- result = {
- 'r2_score': float(r2),
- 'n_components': int(self.model.n_components),
- 'feature_importance': None
- }
- logging.info(f"Final result: {result}")
- logging.info("=== PLS Training End ===")
-
- # 保存训练好的模型
- equipment_type = 'missile' # 或者从参数中获取
- self.save_model(equipment_type)
-
- return result
-
- except Exception as e:
- logging.error(f"Error in PLS training: {str(e)}")
- logging.error(f"Error traceback:", exc_info=True)
- raise Exception(f"PLS training error: {str(e)}")
-
- def _calculate_confidence_intervals(self, predictions, confidence=0.95):
- """
- 计算预测值的置信区间
- """
- try:
- # 使用 bootstrap 方法计算置信区间
- n_predictions = 1000
- bootstrap_predictions = []
-
- for _ in range(n_predictions):
- # 添加随机噪声
- noise = np.random.normal(0, predictions.mean() * 0.05, predictions.shape)
- noisy_pred = predictions + noise
- bootstrap_predictions.append(noisy_pred)
-
- bootstrap_predictions = np.array(bootstrap_predictions).flatten()
-
- # 计算置信区间
- lower = np.percentile(bootstrap_predictions, ((1 - confidence) / 2) * 100)
- upper = np.percentile(bootstrap_predictions, (1 - (1 - confidence) / 2) * 100)
-
- return {
- 'lower': float(lower),
- 'upper': float(upper)
- }
-
- except Exception as e:
- logging.error(f"Error calculating confidence intervals: {str(e)}")
- # 如果计算失败,返回基于10%标准差的区间
- mean_pred = np.mean(predictions)
- return {
- 'lower': float(mean_pred * 0.9),
- 'upper': float(mean_pred * 1.1)
- }
-
- def _get_feature_importance(self):
- """
- 计算特征重要性
- """
- try:
- if not hasattr(self.model, 'x_weights_'):
- return {}
-
- # 获取 VIP 分数
- t = self.model.x_scores_
- w = self.model.x_weights_
- q = self.model.y_loadings_
-
- # 计算每个特征的 VIP 分数
- m, p = w.shape
- vips = np.zeros((p,))
-
- s = np.diag(t.T @ t @ q.T @ q).reshape(m, -1)
- total_s = np.sum(s)
-
- for i in range(p):
- weight = np.array([(w[j,i] / np.linalg.norm(w[:,i]))**2 for j in range(m)])
- vips[i] = np.sqrt(p*(s.T @ weight)/total_s)
-
- # 创建特征重要性字典
- feature_importance = {}
- for i, score in enumerate(vips):
- feature_name = f"feature_{i}" if self.feature_names is None else self.feature_names[i]
- feature_importance[feature_name] = float(score)
-
- # 按重要性排序
- return dict(sorted(feature_importance.items(), key=lambda x: x[1], reverse=True))
-
- except Exception as e:
- logging.error(f"Error calculating feature importance: {str(e)}")
- return {}
-
- def save_model(self, equipment_type):
- """
- 保存模型和标准化器
- """
- try:
- timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
- model_dir = 'models'
- os.makedirs(model_dir, exist_ok=True)
-
- # 保存模型文件
- model_path = f'{model_dir}/pls_{equipment_type}_{timestamp}'
- joblib.dump({
- 'model': self.model,
- 'scaler_X': self.scaler_X,
- 'scaler_y': self.scaler_y,
- 'feature_names': self.feature_names
- }, f'{model_path}.joblib')
-
- # 更新数据库中的模型记录
- with get_db_connection() as conn:
- cursor = conn.cursor()
-
- # 将之前的激活模型设置为非激活
- cursor.execute("""
- UPDATE trained_models
- SET is_active = FALSE
- WHERE equipment_type = %s AND model_type = 'pls'
- """, (equipment_type,))
-
- # 插入新的模型记录
- cursor.execute("""
- INSERT INTO trained_models (
- model_name, model_type, equipment_type, model_path,
- r2_score, training_date, is_active, created_by
- ) VALUES (%s, %s, %s, %s, %s, NOW(), TRUE, 'system')
- """, (
- f'PLS_{timestamp}',
- 'pls',
- equipment_type,
- f'{model_path}.joblib',
- float(self.r2_score_)
- ))
-
- conn.commit()
-
- self.model_path = f'{model_path}.joblib'
- logging.info(f"Model saved to {self.model_path}")
-
- except Exception as e:
- logging.error(f"Error saving model: {str(e)}")
- raise Exception(f"Failed to save model: {str(e)}")
-
- def load_model(self):
- """
- 加载最新的激活模型
- """
- try:
- with get_db_connection() as conn:
- cursor = conn.cursor(dictionary=True)
-
- # 获取最新的激活模型
- cursor.execute("""
- SELECT * FROM trained_models
- WHERE model_type = 'pls' AND is_active = TRUE
- ORDER BY training_date DESC LIMIT 1
- """)
-
- model_record = cursor.fetchone()
-
- if model_record and os.path.exists(model_record['model_path']):
- # 加载模型文件
- saved_data = joblib.load(model_record['model_path'])
- self.model = saved_data['model']
- self.scaler_X = saved_data['scaler_X']
- self.scaler_y = saved_data['scaler_y']
- self.feature_names = saved_data['feature_names']
- self.model_path = model_record['model_path']
-
- logging.info(f"Loaded model from {self.model_path}")
- return True
-
- return False
-
- except Exception as e:
- logging.error(f"Error loading model: {str(e)}")
- return False
\ No newline at end of file
diff --git a/src/routes.py b/src/routes.py
index 1702c22..aff7329 100644
--- a/src/routes.py
+++ b/src/routes.py
@@ -8,22 +8,18 @@ import numpy as np
import mysql.connector
from sklearn.metrics import mean_absolute_error
from .create_template import create_excel_template
-from .pls_regression import PLSPredictor
import json
import os
import time
from .data_preparation import DataPreparation
from .model_trainer import ModelTrainer
+from .logger import setup_logger
# 创建蓝图
api_bp = Blueprint('api', __name__)
-# 配置日志
-logging.basicConfig(
- filename='logs/api.log',
- level=logging.INFO,
- format='%(asctime)s - %(levelname)s - %(message)s'
-)
+# 获取logger
+logger = setup_logger(__name__)
@api_bp.route('/', methods=['GET'])
def index():
@@ -65,44 +61,43 @@ def predict_cost():
"""
try:
data = request.get_json()
-
- # 记录请求
- logging.info(f"Received prediction request for equipment type: {data.get('type', 'unknown')}")
- logging.debug(f"Request data: {data}") # 添加详细的请求数据日志
+ logger.info(f"Received prediction request for equipment type: {data.get('type')}")
# 验证装备类型
if 'type' not in data:
return jsonify({'error': 'Equipment type is required'}), 400
- # 根据装备类型验证必要参数
- required_params = get_required_params(data['type'])
-
- for param in required_params:
- if param not in data:
- return jsonify({'error': f'Missing parameter: {param}'}), 400
-
- # 预���成本
+ # 预测成本
predictor = CostPredictor()
result = predictor.predict(data)
- # 记录预测结果
- logging.info(f"Prediction completed: {result['predicted_cost']}")
-
- # 确保返回的数据格式正确
- response = {
- 'predicted_cost': float(result['predicted_cost']),
- 'confidence_interval': {
- 'lower': float(result['confidence_interval']['lower']),
- 'upper': float(result['confidence_interval']['upper'])
+ # 获取当前使用的模型信息
+ with get_db_connection() as conn:
+ cursor = conn.cursor(dictionary=True)
+ cursor.execute("""
+ SELECT model_type, model_name, r2_score, mae, rmse
+ FROM trained_models
+ WHERE equipment_type = %s AND model_type != 'pls' AND is_active = TRUE
+ LIMIT 1
+ """, (data['type'],))
+ model_info = cursor.fetchone()
+
+ # 在结果中添加模型信息
+ result.update({
+ 'model_info': {
+ 'type': model_info['model_type'],
+ 'name': model_info['model_name'],
+ 'r2_score': float(model_info['r2_score']),
+ 'mae': float(model_info['mae']),
+ 'rmse': float(model_info['rmse'])
}
- }
+ })
- logging.info(f"Sending response: {response}")
- return jsonify(response)
+ logger.info(f"Prediction completed: {result}")
+ return jsonify(result)
except Exception as e:
- logging.error(f"Error in prediction: {str(e)}")
- logging.exception("Detailed error traceback:")
+ logger.error(f"Error in prediction: {str(e)}")
return jsonify({'error': str(e)}), 500
@api_bp.route('/analyze-features', methods=['POST'])
@@ -114,10 +109,10 @@ def analyze_features():
data = request.get_json()
dataset_id = data.get('dataset_id')
- logging.info(f"Starting feature analysis for dataset {dataset_id}")
+ logger.info(f"Starting feature analysis for dataset {dataset_id}")
if not dataset_id:
- logging.warning("No dataset_id provided")
+ logger.warning("No dataset_id provided")
return jsonify({'error': '请选择数据集'}), 400
with get_db_connection() as conn:
@@ -136,10 +131,10 @@ def analyze_features():
dataset = cursor.fetchone()
if not dataset:
- logging.warning(f"Dataset {dataset_id} not found")
+ logger.warning(f"Dataset {dataset_id} not found")
return jsonify({'error': '数据集不存在'}), 404
- logging.info(f"Dataset info: {dataset}")
+ logger.info(f"Dataset info: {dataset}")
# 创建特征分析实例
from src.feature_analysis import FeatureAnalysis
@@ -147,7 +142,7 @@ def analyze_features():
# 获取特征列表
feature_names = analyzer.get_equipment_specific_features(dataset['equipment_type'])
- logging.info(f"Feature names: {feature_names}")
+ logger.info(f"Feature names: {feature_names}")
# 获取数据集中的装备数据
if dataset['equipment_type'] == '火箭炮':
@@ -174,10 +169,10 @@ def analyze_features():
""", (dataset_id,))
equipment_data = cursor.fetchall()
- logging.info(f"Found {len(equipment_data)} equipment records")
+ logger.info(f"Found {len(equipment_data)} equipment records")
if not equipment_data:
- logging.warning("No valid equipment data found in dataset")
+ logger.warning("No valid equipment data found in dataset")
return jsonify({'error': '数据集没有有效的成本数据'}), 400
# 统计每个特征的缺失率
@@ -186,11 +181,11 @@ def analyze_features():
missing_count = sum(1 for item in equipment_data if item.get(name) is None)
missing_rate = missing_count / len(equipment_data)
missing_rates[name] = missing_rate
- logging.info(f"Feature {name} missing rate: {missing_rate:.2%}")
+ logger.info(f"Feature {name} missing rate: {missing_rate:.2%}")
# 过滤掉缺失率过高的特征
valid_features = [name for name in feature_names if missing_rates[name] < 0.7]
- logging.info(f"Valid features after filtering: {valid_features}")
+ logger.info(f"Valid features after filtering: {valid_features}")
if len(valid_features) < 3: # 至少需要3个特征
return jsonify({'error': '有效特征数量不足'}), 400
@@ -200,7 +195,7 @@ def analyze_features():
for name in valid_features:
values = [float(item[name]) for item in equipment_data if item.get(name) is not None]
feature_means[name] = sum(values) / len(values) if values else 0
- logging.info(f"Feature {name} mean value: {feature_means[name]:.2f}")
+ logger.info(f"Feature {name} mean value: {feature_means[name]:.2f}")
# 准备特征和目标值
features = []
@@ -215,32 +210,32 @@ def analyze_features():
# 确保数值类型转换正确
feature_values.append(float(value) if value is not None else feature_means[name])
except (ValueError, TypeError) as e:
- logging.error(f"Error converting value for feature {name}: {value}")
- logging.error(f"Error details: {str(e)}")
+ logger.error(f"Error converting value for feature {name}: {value}")
+ logger.error(f"Error details: {str(e)}")
return jsonify({'error': f'特征 {name} 的值 {value} 无法转换为数值'}), 400
features.append(feature_values)
- # 确保成本值是数值类型
+ # 确保成本值是值类型
try:
target.append(float(item['actual_cost']))
except (ValueError, TypeError) as e:
- logging.error(f"Error converting actual_cost: {item['actual_cost']}")
- logging.error(f"Error details: {str(e)}")
+ logger.error(f"Error converting actual_cost: {item['actual_cost']}")
+ logger.error(f"Error details: {str(e)}")
return jsonify({'error': '成本值无法换为数值'}), 400
- logging.info(f"Prepared {len(features)} feature vectors")
- logging.info(f"First feature vector: {features[0] if features else None}")
- logging.info(f"First target value: {target[0] if target else None}")
+ logger.info(f"Prepared {len(features)} feature vectors")
+ logger.info(f"First feature vector: {features[0] if features else None}")
+ logger.info(f"First target value: {target[0] if target else None}")
# 调用特征分析方法
result = analyzer.analyze_features(features, target, valid_features)
- logging.info("Analysis completed successfully")
+ logger.info("Analysis completed successfully")
return jsonify(result)
except Exception as e:
- logging.error(f"Error analyzing features: {str(e)}")
- logging.error("Detailed traceback:", exc_info=True)
+ logger.error(f"Error analyzing features: {str(e)}")
+ logger.error("Detailed traceback:", exc_info=True)
return jsonify({'error': str(e)}), 500
@api_bp.route('/train', methods=['POST'])
@@ -250,15 +245,15 @@ def train_model():
"""
try:
data = request.get_json()
+ logger.info(f"Starting model training for {data.get('type')}")
equipment_type = data.get('type')
train_dataset_id = data.get('train_dataset_id')
validation_dataset_id = data.get('validation_dataset_id')
models = data.get('models', [])
- logging.info(f"Starting model training for {equipment_type}")
- logging.info(f"Training dataset: {train_dataset_id}")
- logging.info(f"Validation dataset: {validation_dataset_id}")
- logging.info(f"Selected models: {models}")
+ logger.info(f"Training dataset: {train_dataset_id}")
+ logger.info(f"Validation dataset: {validation_dataset_id}")
+ logger.info(f"Selected models: {models}")
# 获取训练数据
with get_db_connection() as conn:
@@ -357,8 +352,8 @@ def train_model():
return jsonify(training_result)
except Exception as e:
- logging.error(f"Error in model training: {str(e)}")
- logging.error("Detailed traceback:", exc_info=True)
+ logger.error(f"Error in model training: {str(e)}")
+ logger.error("Detailed traceback:", exc_info=True)
return jsonify({'error': str(e)}), 500
@api_bp.route('/evaluate', methods=['POST'])
@@ -368,7 +363,7 @@ def evaluate_model():
"""
try:
data = request.get_json()
- logging.info("Received model evaluation request")
+ logger.info("Received model evaluation request")
if 'test_data' not in data:
return jsonify({'error': 'Test data is required'}), 400
@@ -379,11 +374,11 @@ def evaluate_model():
data['test_data']['predicted']
)
- logging.info("Model evaluation completed")
+ logger.info("Model evaluation completed")
return jsonify(evaluation_result)
except Exception as e:
- logging.error(f"Error in model evaluation: {str(e)}")
+ logger.error(f"Error in model evaluation: {str(e)}")
return jsonify({'error': str(e)}), 500
def get_required_params(equipment_type):
@@ -424,7 +419,7 @@ def not_found(error):
@api_bp.errorhandler(500)
def internal_error(error):
- logging.error(f"Internal server error: {str(error)}")
+ logger.error(f"Internal server error: {str(error)}")
return jsonify({'error': 'Internal server error'}), 500
@api_bp.route('/data', methods=['GET'])
@@ -446,10 +441,10 @@ def get_equipment_data():
LIMIT 5
""")
test_params = cursor.fetchall()
- logging.info(f"Test custom params: {test_params}")
+ logger.info(f"Test custom params: {test_params}")
# 获取火箭炮数据
- logging.info("Fetching rocket artillery data...")
+ logger.info("Fetching rocket artillery data...")
cursor.execute("""
SELECT
e.id,
@@ -503,13 +498,13 @@ def get_equipment_data():
WHERE e.type = '火箭炮'
""")
rocket_artillery = cursor.fetchall()
- logging.info(f"Found {len(rocket_artillery)} rocket artillery records")
+ logger.info(f"Found {len(rocket_artillery)} rocket artillery records")
if rocket_artillery:
- logging.info(f"First rocket artillery: {rocket_artillery[0]['name']}")
- logging.info(f"First rocket custom_params: {rocket_artillery[0].get('custom_params')}")
+ logger.info(f"First rocket artillery: {rocket_artillery[0]['name']}")
+ logger.info(f"First rocket custom_params: {rocket_artillery[0].get('custom_params')}")
# 获取巡飞弹数据
- logging.info("Fetching missile data...")
+ logger.info("Fetching missile data...")
cursor.execute("""
SELECT
e.id,
@@ -560,28 +555,28 @@ def get_equipment_data():
WHERE e.type = '巡飞弹'
""")
loitering_munition = cursor.fetchall()
- logging.info(f"Found {len(loitering_munition)} missile records")
+ logger.info(f"Found {len(loitering_munition)} missile records")
if loitering_munition:
- logging.info(f"First missile: {loitering_munition[0]['name']}")
- logging.info(f"First missile custom_params: {loitering_munition[0].get('custom_params')}")
+ logger.info(f"First missile: {loitering_munition[0]['name']}")
+ logger.info(f"First missile custom_params: {loitering_munition[0].get('custom_params')}")
- # 处理 custom_params,���保不为 NULL
+ # 处理 custom_params,保为 NULL
for item in rocket_artillery + loitering_munition:
if item['custom_params'] is None:
item['custom_params'] = []
- logging.debug(f"Set empty custom_params for equipment {item['id']}")
+ logger.debug(f"Set empty custom_params for equipment {item['id']}")
else:
- logging.debug(f"Equipment {item['id']} has {len(item['custom_params'])} custom params")
+ logger.debug(f"Equipment {item['id']} has {len(item['custom_params'])} custom params")
- logging.info("Data fetching completed")
+ logger.info("Data fetching completed")
return jsonify({
'rocket_artillery': rocket_artillery,
'loitering_munition': loitering_munition
})
except Exception as e:
- logging.error(f"Error getting equipment data: {str(e)}")
- logging.error("Detailed traceback:", exc_info=True)
+ logger.error(f"Error getting equipment data: {str(e)}")
+ logger.error("Detailed traceback:", exc_info=True)
return jsonify({'error': str(e)}), 500
@api_bp.route('/data/', methods=['DELETE'])
@@ -607,7 +602,7 @@ def delete_equipment(id):
return jsonify({'status': 'success'})
except Exception as e:
- logging.error(f"Error deleting equipment: {str(e)}")
+ logger.error(f"Error deleting equipment: {str(e)}")
return jsonify({'error': str(e)}), 500
@api_bp.route('/data/template', methods=['GET'])
@@ -633,7 +628,7 @@ def download_template():
)
except Exception as e:
- logging.error(f"Error creating template: {str(e)}")
+ logger.error(f"Error creating template: {str(e)}")
return jsonify({'error': str(e)}), 500
def get_db_connection():
@@ -654,115 +649,60 @@ def pls_predict():
"""
try:
data = request.get_json()
-
- # 记录请求
- logging.info(f"Received PLS prediction request for equipment type: {data.get('type', 'unknown')}")
- logging.debug(f"Request data: {data}")
+ logger.info(f"Received PLS prediction request for equipment type: {data.get('type')}")
# 验证装备类型
if 'type' not in data:
return jsonify({'error': 'Equipment type is required'}), 400
- # 创建PLS预测器
- predictor = PLSPredictor()
- result = predictor.predict(data)
+ # 使用 ModelTrainer 中的 PLS 模型进行预测
+ trainer = ModelTrainer()
+ if not trainer.load_model(data['type'], model_type='pls'): # 指定加载 PLS 模型
+ return jsonify({'error': '未找到可用的模型'}), 404
+
+ # 准备特征数据
+ feature_analyzer = FeatureAnalysis()
+ features = feature_analyzer.get_equipment_specific_features(data['type'])
+ X = np.array([[data.get(feature) for feature in features]])
+ # 预测
+ result = trainer.predict(X)
+
+ # 计算置信区间
+ confidence_interval = trainer._calculate_confidence_interval(result[0])
+
+ # 获取模型信息
+ with get_db_connection() as conn:
+ cursor = conn.cursor(dictionary=True)
+ cursor.execute("""
+ SELECT model_type, model_name, r2_score, mae, rmse
+ FROM trained_models
+ WHERE equipment_type = %s AND model_type = 'pls' AND is_active = TRUE
+ LIMIT 1
+ """, (data['type'],))
+ model_info = cursor.fetchone()
+
# 确保返回的数据可以序列化为JSON
response = {
- 'predicted_cost': float(result['predicted_cost']),
+ 'predicted_cost': float(result[0]),
+ 'model_info': {
+ 'type': model_info['model_type'],
+ 'name': model_info['model_name'],
+ 'r2_score': model_info['r2_score'],
+ 'mae': model_info['mae'],
+ 'rmse': model_info['rmse']
+ },
'confidence_interval': {
- 'lower': float(result['confidence_interval']['lower']),
- 'upper': float(result['confidence_interval']['upper'])
+ 'lower': float(confidence_interval[0]),
+ 'upper': float(confidence_interval[1])
}
}
- # 如果有特征重要性数据也进行转换
- if 'feature_importance' in result:
- response['feature_importance'] = {
- k: float(v) for k, v in result['feature_importance'].items()
- }
-
- logging.info(f"PLS prediction completed: {response}")
+ logger.info(f"PLS prediction completed: {response}")
return jsonify(response)
except Exception as e:
- logging.error(f"Error in PLS prediction: {str(e)}")
- return jsonify({'error': str(e)}), 500
-
-@api_bp.route('/pls/train', methods=['POST'])
-def pls_train():
- """
- PLS模型训练接口
- """
- try:
- # 检查请求类型
- if request.content_type and 'multipart/form-data' in request.content_type:
- # 处理文件上传
- if 'file' not in request.files:
- return jsonify({'error': '没有上传文件'}), 400
-
- file = request.files['file']
- if not file.filename.endswith(('.xls', '.xlsx')):
- return jsonify({'error': '请上传Excel文件'}), 400
-
- # 读取Excel文件
- df = pd.read_excel(file, sheet_name='火箭炮基本参数')
- logging.info(f"Excel data columns: {df.columns}")
-
- # 获取数值列
- numeric_features = df.select_dtypes(include=[np.number]).columns.tolist()
-
- # 检查是否存在成本列
- cost_column = None
- for col in df.columns:
- if '成本' in col or 'cost' in col.lower():
- cost_column = col
- numeric_features.remove(col)
- break
-
- if not cost_column:
- raise ValueError("Excel文件中未找到成本列")
-
- # 准备训练数据
- X = df[numeric_features].values
- y = df[cost_column].values
-
- else:
- # 处理JSON数
- data = request.get_json()
- logging.info(f"Received PLS training data: {data}")
-
- # 将训练数据转换为DataFrame
- training_data = pd.DataFrame(data['training_data'])
-
- # 取值列
- numeric_features = training_data.select_dtypes(include=[np.number]).columns.tolist()
-
- # 准备特征矩阵X和目标变量y
- X = training_data[numeric_features].values
- y = np.array(data['actual_costs'])
-
- logging.info(f"X shape: {X.shape}")
- logging.info(f"y shape: {y.shape}")
-
- # 创建并训练PLS预测器
- predictor = PLSPredictor()
- result = predictor.fit(X, y)
-
- # 确保返回的数据可以序列化为JSON
- response = {
- 'r2_score': float(result['r2_score']),
- 'n_components': int(result['n_components']),
- 'feature_importance': {
- str(k): float(v) for k, v in result['feature_importance'].items()
- } if result['feature_importance'] else {}
- }
-
- logging.info(f"Training completed: {response}")
- return jsonify(response)
-
- except Exception as e:
- logging.error(f"Error in PLS training: {str(e)}")
+ logger.error(f"Error in PLS prediction: {str(e)}")
return jsonify({'error': str(e)}), 500
@api_bp.route('/data/import', methods=['POST'])
@@ -794,7 +734,7 @@ def import_data():
})
except Exception as e:
- logging.error(f"Error importing data: {str(e)}")
+ logger.error(f"Error importing data: {str(e)}")
return jsonify({'error': str(e)}), 500
@api_bp.route('/data/', methods=['PUT'])
@@ -804,8 +744,8 @@ def update_equipment(id):
"""
try:
data = request.get_json()
- logging.info(f"Updating equipment ID: {id}")
- logging.info(f"Update data: {data}")
+ logger.info(f"Updating equipment ID: {id}")
+ logger.info(f"Update data: {data}")
with get_db_connection() as conn:
cursor = conn.cursor()
@@ -816,7 +756,7 @@ def update_equipment(id):
SET name = %s, manufacturer = %s
WHERE id = %s
""", (data['name'], data['manufacturer'], id))
- logging.info("Basic info updated")
+ logger.info("Basic info updated")
# 更新通用参数
cursor.execute("""
@@ -828,7 +768,7 @@ def update_equipment(id):
data['length_m'], data['width_m'], data['height_m'],
data['weight_kg'], data['max_range_km'], id
))
- logging.info("Common params updated")
+ logger.info("Common params updated")
# 根据备类型更新特有参数
if data['type'] == '火箭炮':
@@ -843,7 +783,7 @@ def update_equipment(id):
data['rocket_length_m'], data['rocket_diameter_mm'],
data['rocket_weight_kg'], data['rate_of_fire'], id
))
- logging.info("Rocket artillery params updated")
+ logger.info("Rocket artillery params updated")
else:
cursor.execute("""
UPDATE loitering_munition_params
@@ -858,7 +798,7 @@ def update_equipment(id):
data['launch_mode'], data['folded_length_mm'],
data['folded_width_mm'], data['folded_height_mm'], id
))
- logging.info("Missile params updated")
+ logger.info("Missile params updated")
# 更新成本数据
if 'actual_cost' in data:
@@ -867,27 +807,27 @@ def update_equipment(id):
SET actual_cost = %s
WHERE equipment_id = %s
""", (data['actual_cost'], id))
- logging.info("Cost data updated")
+ logger.info("Cost data updated")
# 更新特殊参数
if 'custom_params' in data and data['custom_params']:
- logging.info(f"Updating custom params: {data['custom_params']}")
+ logger.info(f"Updating custom params: {data['custom_params']}")
for param in data['custom_params']:
cursor.execute("""
UPDATE custom_params
SET param_value = %s
WHERE id = %s AND equipment_id = %s
""", (param['param_value'], param['id'], id))
- logging.info("Custom params updated")
+ logger.info("Custom params updated")
conn.commit()
- logging.info("All updates committed successfully")
+ logger.info("All updates committed successfully")
return jsonify({'success': True})
except Exception as e:
- logging.error(f"Error updating equipment: {str(e)}")
- logging.error("Detailed traceback:", exc_info=True)
+ logger.error(f"Error updating equipment: {str(e)}")
+ logger.error("Detailed traceback:", exc_info=True)
return jsonify({'error': str(e)}), 500
@api_bp.route('/data/details/', methods=['GET'])
@@ -896,7 +836,7 @@ def get_equipment_details(id):
获取装备详数据
"""
try:
- logging.info(f"Getting details for equipment ID: {id}")
+ logger.info(f"Getting details for equipment ID: {id}")
with get_db_connection() as conn:
cursor = conn.cursor(dictionary=True)
@@ -906,11 +846,11 @@ def get_equipment_details(id):
equipment = cursor.fetchone()
if not equipment:
- logging.warning(f"Equipment not found: {id}")
+ logger.warning(f"Equipment not found: {id}")
return jsonify({'error': 'Equipment not found'}), 404
equipment_type = equipment['type']
- logging.info(f"Equipment type: {equipment_type}")
+ logger.info(f"Equipment type: {equipment_type}")
# 根据装备类型选择查询
if equipment_type == '火箭炮':
@@ -984,13 +924,13 @@ def get_equipment_details(id):
result = cursor.fetchone()
if result:
- logging.info(f"Found equipment details: {result['name']}")
- logging.info(f"Custom params: {result.get('custom_params')}")
+ logger.info(f"Found equipment details: {result['name']}")
+ logger.info(f"Custom params: {result.get('custom_params')}")
return jsonify(result)
except Exception as e:
- logging.error(f"Error getting equipment details: {str(e)}")
+ logger.error(f"Error getting equipment details: {str(e)}")
return jsonify({'error': str(e)}), 500
# 添加数据集相关的路由
@@ -1022,7 +962,7 @@ def get_datasets():
return jsonify(datasets)
except Exception as e:
- logging.error(f"Error getting datasets: {str(e)}")
+ logger.error(f"Error getting datasets: {str(e)}")
return jsonify({'error': str(e)}), 500
@api_bp.route('/datasets/', methods=['GET'])
@@ -1076,7 +1016,7 @@ def get_dataset(id):
dataset['equipment'] = equipment
return jsonify(dataset)
except Exception as e:
- logging.error(f"Error getting dataset: {str(e)}")
+ logger.error(f"Error getting dataset: {str(e)}")
return jsonify({'error': str(e)}), 500
@api_bp.route('/datasets', methods=['POST'])
@@ -1108,7 +1048,7 @@ def create_dataset():
conn.commit()
return jsonify({'id': dataset_id, 'message': '数据集创建成功'})
except Exception as e:
- logging.error(f"Error creating dataset: {str(e)}")
+ logger.error(f"Error creating dataset: {str(e)}")
return jsonify({'error': str(e)}), 500
@api_bp.route('/datasets/', methods=['PUT'])
@@ -1142,7 +1082,7 @@ def update_dataset(id):
conn.commit()
return jsonify({'success': True})
except Exception as e:
- logging.error(f"Error updating dataset: {str(e)}")
+ logger.error(f"Error updating dataset: {str(e)}")
return jsonify({'error': str(e)}), 500
@api_bp.route('/datasets/', methods=['DELETE'])
@@ -1163,13 +1103,13 @@ def delete_dataset(id):
conn.commit()
return jsonify({'success': True})
except Exception as e:
- logging.error(f"Error deleting dataset: {str(e)}")
+ logger.error(f"Error deleting dataset: {str(e)}")
return jsonify({'error': str(e)}), 500
@api_bp.route('/models//latest', methods=['GET'])
def get_latest_model(equipment_type):
"""
- 获取最新训练的���型信息
+ 获取最新训练的型信息
"""
try:
with get_db_connection() as conn:
@@ -1184,7 +1124,7 @@ def get_latest_model(equipment_type):
return jsonify(model)
except Exception as e:
- logging.error(f"Error getting latest model: {str(e)}")
+ logger.error(f"Error getting latest model: {str(e)}")
return jsonify({'error': str(e)}), 500
@api_bp.route('/models', methods=['GET'])
@@ -1218,7 +1158,7 @@ def get_models():
return jsonify(models)
except Exception as e:
- logging.error(f"Error getting models: {str(e)}")
+ logger.error(f"Error getting models: {str(e)}")
return jsonify({'error': str(e)}), 500
@api_bp.route('/models//activate', methods=['POST'])
@@ -1258,7 +1198,7 @@ def activate_model(id):
return jsonify({'success': True})
except Exception as e:
- logging.error(f"Error activating model: {str(e)}")
+ logger.error(f"Error activating model: {str(e)}")
return jsonify({'error': str(e)}), 500
@api_bp.route('/models/', methods=['DELETE'])
@@ -1294,5 +1234,23 @@ def delete_model(id):
return jsonify({'success': True})
except Exception as e:
- logging.error(f"Error deleting model: {str(e)}")
+ logger.error(f"Error deleting model: {str(e)}")
+ return jsonify({'error': str(e)}), 500
+
+@api_bp.route('/predict/all', methods=['POST'])
+def predict_all():
+ """
+ 获取所有机器学习模型的预测结果
+ """
+ try:
+ data = request.get_json()
+ logger.info(f"Received prediction request for all models, equipment type: {data.get('type')}")
+
+ predictor = CostPredictor()
+ results = predictor.predict_all(data)
+
+ return jsonify(results)
+
+ except Exception as e:
+ logger.error(f"Error in prediction: {str(e)}")
return jsonify({'error': str(e)}), 500
\ No newline at end of file
diff --git a/src/run.py b/src/run.py
deleted file mode 100644
index 73acb7d..0000000
--- a/src/run.py
+++ /dev/null
@@ -1,28 +0,0 @@
-import os
-import logging
-from src.app import app
-
-# 确保必要的目录存在
-os.makedirs('logs', exist_ok=True)
-os.makedirs('models', exist_ok=True)
-os.makedirs('data', exist_ok=True)
-
-# 配置日志
-logging.basicConfig(
- filename='logs/server.log',
- level=logging.INFO,
- format='%(asctime)s - %(levelname)s - %(message)s'
-)
-
-if __name__ == "__main__":
- try:
- logging.info("Starting server...")
- app.run(
- host='localhost',
- port=5001,
- debug=True, # 启用调试模式
- use_reloader=True # 启用自动重载
- )
- except Exception as e:
- logging.error(f"Server failed to start: {str(e)}")
- raise
\ No newline at end of file