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) }} 元

+
+ + + + + + - + - + - + - - + + + + - + - + @@ -106,148 +115,142 @@

特征重要性

- - - + + +
- - -
-

最佳模型

- - - {{ 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