完成了基本功能

This commit is contained in:
Tian jianyong 2024-11-09 16:48:50 +08:00
parent 6e5172962a
commit fccd4c4366
28 changed files with 1211 additions and 1425 deletions

View File

@ -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.

8
app.py
View File

@ -1,8 +0,0 @@
import logging
# 配置日志
logging.basicConfig(
filename='logs/api.log',
level=logging.INFO,
format='%(asctime)s - %(levelname)s - %(message)s'
)

View File

@ -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%

View File

@ -16,5 +16,5 @@
"scripthost"
]
},
"exclude": ["**/HelloWorld.vue"]
"include": ["src/**/*"]
}

View File

@ -7,13 +7,12 @@
:default-active="$route.path"
>
<el-menu-item index="/">首页</el-menu-item>
<el-menu-item index="/predict">机器学习预测</el-menu-item>
<el-menu-item index="/pls-predict">PLS回归预测</el-menu-item>
<el-menu-item index="/predict">成本预测</el-menu-item>
<el-menu-item index="/analysis">特征分析</el-menu-item>
<el-menu-item index="/training">模型训练</el-menu-item>
<el-menu-item index="/data">数据管理</el-menu-item>
<el-menu-item index="/datasets">数据集管理</el-menu-item>
<el-menu-item index="/models">模型管理</el-menu-item>
<el-menu-item index="/datasets">数据集管理</el-menu-item>
<el-menu-item index="/data">数据管理</el-menu-item>
</el-menu>
</el-header>

View File

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

View File

@ -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',

View File

@ -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: '',

View File

@ -647,7 +647,7 @@ const isNumericParam = (param) => {
}
// handleTabClick
const handleTabClick = (tab) => {
const handleTabClick = () => {
//
searchQuery.value = ''
filterManufacturer.value = ''

View File

@ -8,15 +8,8 @@
<el-col :span="8">
<el-card @click="$router.push('/predict')">
<el-icon><Money /></el-icon>
<h3>机器学习预测</h3>
<p>基于机器学习模型的成本预测</p>
</el-card>
</el-col>
<el-col :span="8">
<el-card @click="$router.push('/pls-predict')">
<el-icon><TrendCharts /></el-icon>
<h3>PLS回归预测</h3>
<p>基于偏最小二乘回归的成本预测</p>
<h3>成本预测</h3>
<p>基于机器学习和 PLS 回归模型的成本预测</p>
</el-card>
</el-col>
<el-col :span="8">
@ -34,10 +27,10 @@
</el-card>
</el-col>
<el-col :span="8">
<el-card @click="$router.push('/data')">
<el-card @click="$router.push('/models')">
<el-icon><Management /></el-icon>
<h3>数据管理</h3>
<p>管理装备数据和成本数据</p>
<h3>模型管理</h3>
<p>管理训练好的模型</p>
</el-card>
</el-col>
<el-col :span="8">
@ -47,13 +40,20 @@
<p>管理训练和验证数据集</p>
</el-card>
</el-col>
<el-col :span="8">
<el-card @click="$router.push('/data')">
<el-icon><Management /></el-icon>
<h3>数据管理</h3>
<p>管理装备数据和成本数据</p>
</el-card>
</el-col>
</el-row>
</el-card>
</div>
</template>
<script setup>
import { Money, DataAnalysis, Monitor, Management, TrendCharts, Collection } from '@element-plus/icons-vue'
import { Money, DataAnalysis, Monitor, Management, Collection } from '@element-plus/icons-vue'
</script>
<style lang="scss" scoped>

View File

@ -8,13 +8,13 @@
</template>
<!-- 模型列表 -->
<el-table :data="models" border style="width: 100%">
<el-table-column prop="model_name" label="模型名称"></el-table-column>
<el-table :data="modelList" border style="width: 100%">
<el-table-column prop="model_type" label="模型类型">
<template #default="scope">
{{ formatModelType(scope.row.model_type) }}
{{ getModelName(scope.row.model_type) }}
</template>
</el-table-column>
<el-table-column prop="model_name" label="模型名称"></el-table-column>
<el-table-column prop="equipment_type" label="装备类型"></el-table-column>
<el-table-column prop="r2_score" label="R²分数">
<template #default="scope">
@ -114,7 +114,7 @@ import { API_BASE_URL } from '@/config'
import * as echarts from 'echarts'
//
const models = ref([])
const modelList = ref([])
const selectedModel = ref(null)
const detailsVisible = ref(false)
const importanceChartRef = ref(null)
@ -124,7 +124,7 @@ const importanceChart = ref(null)
const loadModels = async () => {
try {
const response = await axios.get(`${API_BASE_URL}/models`)
models.value = response.data
modelList.value = response.data
} catch (error) {
ElMessage.error('获取模型列表失败')
}
@ -243,6 +243,17 @@ onUnmounted(() => {
onMounted(() => {
loadModels()
})
const getModelName = (modelType) => {
const modelNames = {
'pls': 'PLS回归',
'xgboost': 'XGBoost',
'lightgbm': 'LightGBM',
'gbm': 'GBM',
'rf': 'Random Forest'
}
return modelNames[modelType] || modelType
}
</script>
<style lang="scss" scoped>

View File

@ -1,243 +0,0 @@
<template>
<div class="predict-page">
<el-card class="predict-card">
<template #header>
<h2>PLS回归预测</h2>
</template>
<el-form :model="formData" label-width="120px">
<!-- 装备类型选择 -->
<el-form-item label="装备类型">
<el-select v-model="formData.type" @change="handleTypeChange">
<el-option label="火箭炮" value="火箭炮"></el-option>
<el-option label="巡飞弹" value="巡飞弹"></el-option>
</el-select>
</el-form-item>
<!-- 通用参数 -->
<el-form-item label="总长(m)">
<el-input-number v-model="formData.length_m" :precision="2"></el-input-number>
</el-form-item>
<el-form-item label="宽度(m)">
<el-input-number v-model="formData.width_m" :precision="2"></el-input-number>
</el-form-item>
<el-form-item label="高度(m)">
<el-input-number v-model="formData.height_m" :precision="2"></el-input-number>
</el-form-item>
<el-form-item label="重量(kg)">
<el-input-number v-model="formData.weight_kg"></el-input-number>
</el-form-item>
<el-form-item label="最大射程(km)">
<el-input-number v-model="formData.max_range_km"></el-input-number>
</el-form-item>
<!-- 火箭炮特有参数 -->
<template v-if="formData.type === '火箭炮'">
<el-form-item label="方向射界(度)">
<el-input-number v-model="formData.firing_angle_horizontal"></el-input-number>
</el-form-item>
<el-form-item label="高低射界(度)">
<el-input-number v-model="formData.firing_angle_vertical"></el-input-number>
</el-form-item>
<el-form-item label="火箭弹长度(m)">
<el-input-number v-model="formData.rocket_length_m" :precision="2"></el-input-number>
</el-form-item>
<el-form-item label="弹体直径(mm)">
<el-input-number v-model="formData.rocket_diameter_mm"></el-input-number>
</el-form-item>
<el-form-item label="火箭弹重量(kg)">
<el-input-number v-model="formData.rocket_weight_kg"></el-input-number>
</el-form-item>
<el-form-item label="射速(发/分钟)">
<el-input-number v-model="formData.rate_of_fire"></el-input-number>
</el-form-item>
</template>
<!-- 巡飞弹特有参数 -->
<template v-if="formData.type === '巡飞弹'">
<el-form-item label="最大速度(km/h)">
<el-input-number v-model="formData.max_speed_kmh"></el-input-number>
</el-form-item>
<el-form-item label="巡航速度(km/h)">
<el-input-number v-model="formData.cruise_speed_kmh"></el-input-number>
</el-form-item>
<el-form-item label="巡飞时间(min)">
<el-input-number v-model="formData.flight_time_min"></el-input-number>
</el-form-item>
<el-form-item label="折叠长度(mm)">
<el-input-number v-model="formData.folded_length_mm"></el-input-number>
</el-form-item>
<el-form-item label="折叠宽度(mm)">
<el-input-number v-model="formData.folded_width_mm"></el-input-number>
</el-form-item>
<el-form-item label="折叠高度(mm)">
<el-input-number v-model="formData.folded_height_mm"></el-input-number>
</el-form-item>
</template>
<el-form-item>
<el-button type="primary" @click="submitForm">预测成本</el-button>
<el-button @click="resetForm">重置</el-button>
</el-form-item>
</el-form>
<!-- 预测结果 -->
<div v-if="predictionResult" class="prediction-result">
<h3>预测结果</h3>
<el-descriptions border>
<el-descriptions-item label="预测成本">
{{ formatMoney(predictionResult.predicted_cost) }}
</el-descriptions-item>
<el-descriptions-item label="置信区间">
{{ formatMoney(predictionResult.confidence_interval.lower) }} ~
{{ formatMoney(predictionResult.confidence_interval.upper) }}
</el-descriptions-item>
</el-descriptions>
</div>
</el-card>
</div>
</template>
<script setup>
import { ref, reactive } from 'vue'
import { ElMessage } from 'element-plus'
import axios from 'axios'
import { API_BASE_URL } from '@/config'
const formData = reactive({
type: '',
length_m: null,
width_m: null,
height_m: null,
weight_kg: null,
max_range_km: null
})
const predictionResult = ref(null)
const handleTypeChange = () => {
if (formData.type === '火箭炮') {
Object.assign(formData, {
firing_angle_horizontal: null,
firing_angle_vertical: null,
rocket_length_m: null,
rocket_diameter_mm: null,
rocket_weight_kg: null,
rate_of_fire: null
})
} else if (formData.type === '巡飞弹') {
Object.assign(formData, {
max_speed_kmh: null,
cruise_speed_kmh: null,
flight_time_min: null,
folded_length_mm: null,
folded_width_mm: null,
folded_height_mm: null
})
}
}
const submitForm = async () => {
try {
//
if (!formData.type) {
throw new Error('请选择装备类型')
}
//
const commonFields = ['length_m', 'width_m', 'height_m', 'weight_kg', 'max_range_km']
for (const field of commonFields) {
if (!formData[field]) {
throw new Error(`请输入${formatFieldName(field)}`)
}
}
//
if (formData.type === '火箭炮') {
const rocketFields = [
'firing_angle_horizontal', 'firing_angle_vertical',
'rocket_length_m', 'rocket_diameter_mm', 'rocket_weight_kg', 'rate_of_fire'
]
for (const field of rocketFields) {
if (!formData[field]) {
throw new Error(`请输入${formatFieldName(field)}`)
}
}
} else if (formData.type === '巡飞弹') {
const missileFields = [
'max_speed_kmh', 'cruise_speed_kmh', 'flight_time_min',
'folded_length_mm', 'folded_width_mm', 'folded_height_mm'
]
for (const field of missileFields) {
if (!formData[field]) {
throw new Error(`请输入${formatFieldName(field)}`)
}
}
}
//
const response = await axios.post(`${API_BASE_URL}/pls/predict`, formData)
predictionResult.value = response.data
} catch (error) {
ElMessage.error(error.message || '预测失败')
}
}
const resetForm = () => {
formData.type = ''
formData.length_m = null
formData.width_m = null
formData.height_m = null
formData.weight_kg = null
formData.max_range_km = null
predictionResult.value = null
}
const formatFieldName = (field) => {
const nameMap = {
'length_m': '总长',
'width_m': '宽度',
'height_m': '高度',
'weight_kg': '重量',
'max_range_km': '最大射程',
'firing_angle_horizontal': '方向射界',
'firing_angle_vertical': '高低射界',
'rocket_length_m': '火箭弹长度',
'rocket_diameter_mm': '弹体直径',
'rocket_weight_kg': '火箭弹重量',
'rate_of_fire': '射速',
'max_speed_kmh': '最大速度',
'cruise_speed_kmh': '巡航速度',
'flight_time_min': '巡飞时间',
'folded_length_mm': '折叠长度',
'folded_width_mm': '折叠宽度',
'folded_height_mm': '折叠高度'
}
return nameMap[field] || field
}
//
const formatMoney = (value) => {
return new Intl.NumberFormat('zh-CN', {
minimumFractionDigits: 2,
maximumFractionDigits: 2
}).format(value)
}
</script>
<style scoped>
.predict-page {
padding: 20px;
}
.predict-card {
max-width: 800px;
margin: 0 auto;
}
.prediction-result {
margin-top: 20px;
padding: 20px;
background-color: #f5f7fa;
border-radius: 4px;
}
</style>

View File

@ -2,11 +2,11 @@
<div class="predict-page">
<el-card class="predict-card">
<template #header>
<h2>装备成本预测</h2>
<h2>成本预测</h2>
</template>
<!-- 装备类型选择 -->
<el-form :model="formData" label-width="120px">
<!-- 装备类型选择 -->
<el-form-item label="装备类型">
<el-select v-model="formData.type" @change="handleTypeChange">
<el-option label="火箭炮" value="火箭炮"></el-option>
@ -95,29 +95,60 @@
</el-form>
<!-- 预测结果 -->
<div v-if="predictionResult" class="prediction-result">
<div v-if="predictionResults" class="prediction-results">
<h3>预测结果</h3>
<el-descriptions border>
<el-descriptions-item label="预测成本">
{{ formatCurrency(predictionResult.predicted_cost) }}
</el-descriptions-item>
<el-descriptions-item label="置信区间">
{{ formatCurrency(predictionResult.confidence_interval.lower) }} ~
{{ formatCurrency(predictionResult.confidence_interval.upper) }}
</el-descriptions-item>
</el-descriptions>
<!-- 机器学习模型预测结果 -->
<div class="ml-prediction">
<h4>机器学习模型预测</h4>
<el-descriptions border>
<el-descriptions-item label="模型类型">
{{ getModelName(mlPrediction.model_info.type) }}
</el-descriptions-item>
<el-descriptions-item label="模型名称">
{{ mlPrediction.model_info.name }}
</el-descriptions-item>
<el-descriptions-item label="预测成本">
{{ formatMoney(mlPrediction.predicted_cost) }}
</el-descriptions-item>
<el-descriptions-item label="置信区间">
{{ formatMoney(mlPrediction.confidence_interval.lower) }} ~
{{ formatMoney(mlPrediction.confidence_interval.upper) }}
</el-descriptions-item>
</el-descriptions>
</div>
<!-- PLS回归预测结果 -->
<div class="pls-prediction">
<h4>PLS回归预测</h4>
<el-descriptions border>
<el-descriptions-item label="模型类型">
{{ getModelName(plsPrediction.model_info.type) }}
</el-descriptions-item>
<el-descriptions-item label="模型名称">
{{ plsPrediction.model_info.name }}
</el-descriptions-item>
<el-descriptions-item label="预测成本">
{{ formatMoney(plsPrediction.predicted_cost) }}
</el-descriptions-item>
<el-descriptions-item label="置信区间">
{{ formatMoney(plsPrediction.confidence_interval.lower) }} ~
{{ formatMoney(plsPrediction.confidence_interval.upper) }}
</el-descriptions-item>
</el-descriptions>
</div>
</div>
</el-card>
</div>
</template>
<script setup>
import { ref } from 'vue'
import { ref, reactive } from 'vue'
import { ElMessage } from 'element-plus'
import axios from 'axios'
import { API_BASE_URL } from '@/config'
const formData = ref({
const formData = reactive({
type: '',
length_m: null,
width_m: null,
@ -126,83 +157,165 @@ const formData = ref({
max_range_km: null
})
const predictionResult = ref(null)
const predictionResults = ref(null)
const mlPrediction = ref(null)
const plsPrediction = ref(null)
const handleTypeChange = () => {
//
if (formData.value.type === '火箭炮') {
formData.value = {
...formData.value,
firing_angle_horizontal: null,
firing_angle_vertical: null,
rocket_length_m: null,
rocket_diameter_mm: null,
rocket_weight_kg: null,
rate_of_fire: null
}
} else if (formData.value.type === '巡飞弹') {
formData.value = {
...formData.value,
max_speed_kmh: null,
cruise_speed_kmh: null,
flight_time_min: null,
warhead_type: '',
launch_mode: '',
folded_length_mm: null,
folded_width_mm: null,
folded_height_mm: null
}
if (formData.type === '火箭炮') {
formData.firing_angle_horizontal = null
formData.firing_angle_vertical = null
formData.rocket_length_m = null
formData.rocket_diameter_mm = null
formData.rocket_weight_kg = null
formData.rate_of_fire = null
} else if (formData.type === '巡飞弹') {
formData.max_speed_kmh = null
formData.cruise_speed_kmh = null
formData.flight_time_min = null
formData.warhead_type = ''
formData.launch_mode = ''
formData.folded_length_mm = null
formData.folded_width_mm = null
formData.folded_height_mm = null
}
}
const submitForm = async () => {
try {
const response = await axios.post(`${API_BASE_URL}/predict`, formData.value)
predictionResult.value = response.data
//
if (!formData.type) {
throw new Error('请选择装备类型')
}
//
const commonFields = ['length_m', 'width_m', 'height_m', 'weight_kg', 'max_range_km']
for (const field of commonFields) {
if (!formData[field]) {
throw new Error(`请输入${formatFieldName(field)}`)
}
}
//
if (formData.type === '火箭炮') {
const rocketFields = [
'firing_angle_horizontal', 'firing_angle_vertical',
'rocket_length_m', 'rocket_diameter_mm', 'rocket_weight_kg', 'rate_of_fire'
]
for (const field of rocketFields) {
if (!formData[field]) {
throw new Error(`请输入${formatFieldName(field)}`)
}
}
} else if (formData.type === '巡飞弹') {
const missileFields = [
'max_speed_kmh', 'cruise_speed_kmh', 'flight_time_min',
'folded_length_mm', 'folded_width_mm', 'folded_height_mm'
]
for (const field of missileFields) {
if (!formData[field]) {
throw new Error(`请输入${formatFieldName(field)}`)
}
}
}
//
const [mlResponse, plsResponse] = await Promise.all([
axios.post(`${API_BASE_URL}/predict`, formData),
axios.post(`${API_BASE_URL}/pls/predict`, formData)
])
mlPrediction.value = mlResponse.data
plsPrediction.value = plsResponse.data
predictionResults.value = true
} catch (error) {
ElMessage.error(error.response?.data?.error || '预测失败')
ElMessage.error(error.message || '预测失败')
}
}
const resetForm = () => {
formData.value = {
type: '',
length_m: null,
width_m: null,
height_m: null,
weight_kg: null,
max_range_km: null
}
predictionResult.value = null
formData.type = ''
formData.length_m = null
formData.width_m = null
formData.height_m = null
formData.weight_kg = null
formData.max_range_km = null
predictionResults.value = null
mlPrediction.value = null
plsPrediction.value = null
}
const formatCurrency = (value) => {
const formatFieldName = (field) => {
const nameMap = {
'length_m': '总长',
'width_m': '宽度',
'height_m': '高度',
'weight_kg': '重量',
'max_range_km': '最大射程',
'firing_angle_horizontal': '方向射界',
'firing_angle_vertical': '高低射界',
'rocket_length_m': '火箭弹长度',
'rocket_diameter_mm': '弹体直径',
'rocket_weight_kg': '火箭弹重量',
'rate_of_fire': '射速',
'max_speed_kmh': '最大速度',
'cruise_speed_kmh': '巡航速度',
'flight_time_min': '巡飞时间',
'folded_length_mm': '折叠长度',
'folded_width_mm': '折叠宽度',
'folded_height_mm': '折叠高度'
}
return nameMap[field] || field
}
const formatMoney = (value) => {
return new Intl.NumberFormat('zh-CN', {
style: 'currency',
currency: 'CNY'
}).format(value)
}
const getModelName = (modelType) => {
const modelNames = {
'pls': 'PLS回归',
'xgboost': 'XGBoost',
'lightgbm': 'LightGBM',
'gbm': 'GBM',
'rf': 'Random Forest'
}
return modelNames[modelType] || modelType
}
</script>
<style lang="scss" scoped>
<style scoped>
.predict-page {
padding: 20px;
}
.predict-card {
max-width: 800px;
margin: 0 auto;
}
.prediction-results {
margin-top: 20px;
.predict-card {
max-width: 800px;
margin: 0 auto;
h2 {
text-align: center;
margin: 0;
}
}
.prediction-result {
.ml-prediction, .pls-prediction {
margin-top: 20px;
padding: 20px;
background-color: #f5f7fa;
border-radius: 4px;
}
h4 {
margin-top: 0;
margin-bottom: 15px;
}
}
.el-descriptions {
margin-top: 10px;
}
</style>

View File

@ -6,52 +6,49 @@
</template>
<!-- 训练配置 -->
<el-form :model="formData" label-width="120px">
<el-form-item label="装备类型" required>
<el-select v-model="formData.type" @change="handleTypeChange">
<el-option label="火箭炮" value="火箭炮"></el-option>
<el-option label="巡飞弹" value="巡飞弹"></el-option>
<el-form :model="trainingConfig" label-width="120px">
<el-form-item label="装备类型">
<el-select v-model="trainingConfig.type" placeholder="选择装备类型">
<el-option label="火箭炮" value="火箭炮" />
<el-option label="巡飞弹" value="巡飞弹" />
</el-select>
</el-form-item>
<!-- 选择训练集 -->
<el-form-item label="训练数据集" required>
<el-select v-model="formData.train_dataset_id" placeholder="选择训练数据集">
<el-option
v-for="dataset in trainingDatasets"
<el-form-item label="训练数据集">
<el-select v-model="trainingConfig.train_dataset_id" placeholder="选择训练数据集">
<el-option
v-for="dataset in trainingDatasets"
:key="dataset.id"
:label="dataset.name"
:value="dataset.id"
></el-option>
/>
</el-select>
</el-form-item>
<!-- 选择验证集 -->
<el-form-item label="验证数据集">
<el-select v-model="formData.validation_dataset_id" placeholder="选择验证数据集" clearable>
<el-option
v-for="dataset in validationDatasets"
<el-select v-model="trainingConfig.validation_dataset_id" placeholder="选择验证数据集">
<el-option
v-for="dataset in validationDatasets"
:key="dataset.id"
:label="dataset.name"
:value="dataset.id"
></el-option>
/>
</el-select>
</el-form-item>
<!-- 模型选择 -->
<el-form-item label="训练模型" required>
<el-checkbox-group v-model="formData.models">
<el-checkbox label="xgboost">XGBoost</el-checkbox>
<el-checkbox label="lightgbm">LightGBM</el-checkbox>
<el-checkbox label="gbdt">GBDT</el-checkbox>
<el-checkbox label="rf">Random Forest</el-checkbox>
<el-form-item label="选择模型">
<el-checkbox-group v-model="trainingConfig.models">
<el-checkbox label="pls" disabled>PLS回归</el-checkbox>
<el-checkbox label="xgboost" checked>XGBoost</el-checkbox>
<el-checkbox label="lightgbm" checked>LightGBM</el-checkbox>
<el-checkbox label="gbm" checked>GBM</el-checkbox>
<el-checkbox label="rf" checked>Random Forest</el-checkbox>
</el-checkbox-group>
</el-form-item>
<!-- 开始训练按钮 -->
<el-form-item>
<el-button type="primary" @click="startTraining" :loading="training">
{{ training ? '训练中...' : '开始训练' }}
<el-button type="primary" @click="startTraining" :loading="isTraining">
开始训练
</el-button>
</el-form-item>
</el-form>
@ -60,44 +57,56 @@
<div v-if="trainingResult" class="training-result">
<h3>训练结果</h3>
<!-- 模型评估指标 -->
<el-table :data="modelMetrics" border style="width: 100%">
<el-table-column prop="model" label="模型">
<!-- 最佳模型信息 -->
<div class="best-model-info" v-if="trainingResult.best_model">
<h4>最佳模型: {{ getModelName(trainingResult.best_model.type) }}</h4>
<p>R²分数: {{ formatNumber(trainingResult.best_model.r2) }}</p>
<p>MAE: {{ formatNumber(trainingResult.best_model.mae) }} </p>
<p>RMSE: {{ formatNumber(trainingResult.best_model.rmse) }} </p>
</div>
<!-- 所有模型评估结果 -->
<el-table :data="modelResults" border style="width: 100%; margin-top: 20px;">
<el-table-column prop="model" label="模型" width="120">
<template #default="scope">
{{ formatModelName(scope.row.model) }}
{{ getModelName(scope.row.model) }}
</template>
</el-table-column>
<!-- 训练集评估 -->
<el-table-column label="训练集评估">
<el-table-column prop="train.r2" label="R²分数">
<el-table-column prop="train.r2" label="R²分数" width="120">
<template #default="scope">
{{ scope.row.train.r2.toFixed(4) }}
{{ formatNumber(scope.row.train.r2) }}
</template>
</el-table-column>
<el-table-column prop="train.mae" label="MAE (元)">
<el-table-column prop="train.mae" label="MAE (元)" width="150">
<template #default="scope">
{{ scope.row.train.mae.toFixed(2) }}
{{ formatNumber(scope.row.train.mae) }}
</template>
</el-table-column>
<el-table-column prop="train.rmse" label="RMSE (元)">
<el-table-column prop="train.rmse" label="RMSE (元)" width="150">
<template #default="scope">
{{ scope.row.train.rmse.toFixed(2) }}
{{ formatNumber(scope.row.train.rmse) }}
</template>
</el-table-column>
</el-table-column>
<el-table-column label="验证集评估" v-if="formData.validation_dataset_id">
<el-table-column prop="validation.r2" label="R²分数">
<!-- 验证集评估 -->
<el-table-column label="验证集评估">
<el-table-column prop="validation.r2" label="R²分数" width="120">
<template #default="scope">
{{ scope.row.validation.r2.toFixed(4) }}
{{ formatNumber(scope.row.validation.r2) }}
</template>
</el-table-column>
<el-table-column prop="validation.mae" label="MAE (元)">
<el-table-column prop="validation.mae" label="MAE (元)" width="150">
<template #default="scope">
{{ scope.row.validation.mae.toFixed(2) }}
{{ formatNumber(scope.row.validation.mae) }}
</template>
</el-table-column>
<el-table-column prop="validation.rmse" label="RMSE (元)">
<el-table-column prop="validation.rmse" label="RMSE (元)" width="150">
<template #default="scope">
{{ scope.row.validation.rmse.toFixed(2) }}
{{ formatNumber(scope.row.validation.rmse) }}
</template>
</el-table-column>
</el-table-column>
@ -106,148 +115,142 @@
<!-- 特征重要性 -->
<div v-if="trainingResult.feature_importance" class="feature-importance">
<h4>特征重要性</h4>
<el-table :data="featureImportanceData" border style="width: 100%">
<el-table-column prop="feature" label="特征"></el-table-column>
<el-table-column prop="importance" label="重要性">
<el-table
:data="featureImportanceData"
border
style="width: 100%; margin-top: 10px;"
>
<el-table-column prop="feature" label="特征" width="180" />
<el-table-column prop="importance" label="重要性" width="120">
<template #default="scope">
<el-progress
:percentage="scope.row.importance * 100"
:format="format => format.toFixed(2) + '%'"
:color="getImportanceColor(scope.row.importance)"
></el-progress>
{{ formatNumber(scope.row.importance) }}
</template>
</el-table-column>
</el-table>
</div>
<!-- 最佳模型信息 -->
<div v-if="trainingResult.best_model" class="best-model">
<h4>最佳模型</h4>
<el-descriptions :column="2" border>
<el-descriptions-item label="模型类型">
{{ formatModelName(trainingResult.best_model.type) }}
</el-descriptions-item>
<el-descriptions-item label="R²分数">
{{ trainingResult.best_model.r2.toFixed(4) }}
</el-descriptions-item>
<el-descriptions-item label="MAE">
{{ formatMoney(trainingResult.best_model.mae) }}
</el-descriptions-item>
<el-descriptions-item label="RMSE">
{{ formatMoney(trainingResult.best_model.rmse) }}
</el-descriptions-item>
</el-descriptions>
</div>
</div>
</el-card>
</div>
</template>
<script setup>
import { ref, computed, onMounted } from 'vue'
import { ref, computed, onMounted, watch } from 'vue'
import { ElMessage } from 'element-plus'
import axios from 'axios'
import { API_BASE_URL } from '@/config'
//
const formData = ref({
//
const trainingConfig = ref({
type: '',
train_dataset_id: null,
validation_dataset_id: null,
models: ['xgboost', 'lightgbm', 'gbdt', 'rf']
models: ['pls']
})
//
const trainingDatasets = ref([])
const validationDatasets = ref([])
const training = ref(false)
//
const isTraining = ref(false)
const trainingResult = ref(null)
//
const loadDatasets = async (type) => {
//
const loadDatasets = async () => {
try {
const response = await axios.get(`${API_BASE_URL}/datasets`, {
params: {
equipment_type: type,
purpose: '训练' //
}
})
trainingDatasets.value = response.data
//
const trainResponse = await axios.get(
`${API_BASE_URL}/datasets`,
{ params: { equipment_type: trainingConfig.value.type, purpose: '训练' } }
)
trainingDatasets.value = trainResponse.data
//
const valResponse = await axios.get(
`${API_BASE_URL}/datasets`,
{ params: { equipment_type: trainingConfig.value.type, purpose: '验证' } }
)
validationDatasets.value = valResponse.data
//
const validationResponse = await axios.get(`${API_BASE_URL}/datasets`, {
params: {
equipment_type: type,
purpose: '验证'
}
})
validationDatasets.value = validationResponse.data
} catch (error) {
ElMessage.error('获取数据集列表失败')
ElMessage.error('加载数据集失败')
console.error('Error loading datasets:', error)
}
}
//
const handleTypeChange = () => {
formData.value.train_dataset_id = null
formData.value.validation_dataset_id = null
loadDatasets(formData.value.type)
}
//
watch(() => trainingConfig.value.type, (newType) => {
if (newType) {
loadDatasets()
}
})
//
const startTraining = async () => {
try {
//
if (!formData.value.type) {
throw new Error('请选择装备类型')
//
if (!trainingConfig.value.type) {
ElMessage.warning('请选择装备类型')
return
}
if (!formData.value.train_dataset_id) {
throw new Error('请选择训练数据集')
if (!trainingConfig.value.train_dataset_id) {
ElMessage.warning('请选择训练数据集')
return
}
if (formData.value.models.length === 0) {
throw new Error('请至少选择一个训练模型')
if (!trainingConfig.value.validation_dataset_id) {
ElMessage.warning('请选择验证数据集')
return
}
if (trainingConfig.value.models.length === 0) {
ElMessage.warning('请至少选择一个模型')
return
}
training.value = true
isTraining.value = true
//
const response = await axios.post(`${API_BASE_URL}/train`, {
type: formData.value.type,
train_dataset_id: formData.value.train_dataset_id,
validation_dataset_id: formData.value.validation_dataset_id,
models: formData.value.models
})
const response = await axios.post(`${API_BASE_URL}/train`, trainingConfig.value)
if (response.data.error) {
throw new Error(response.data.error)
}
trainingResult.value = response.data
ElMessage.success('训练完成')
} catch (error) {
console.error('Training error:', error)
ElMessage.error(error.message || '训练失败')
console.error('Training error:', error)
} finally {
training.value = false
isTraining.value = false
}
}
//
const formatModelName = (name) => {
const nameMap = {
//
const formatNumber = (value) => {
if (value === null || value === undefined) return '-'
if (typeof value === 'number') {
if (Math.abs(value) >= 1000) {
return value.toLocaleString('zh-CN', { maximumFractionDigits: 2 })
}
return value.toFixed(4)
}
return value
}
//
const getModelName = (modelType) => {
const modelNames = {
'xgboost': 'XGBoost',
'lightgbm': 'LightGBM',
'gbdt': 'GBDT',
'gbm': 'GBM',
'rf': 'Random Forest'
}
return nameMap[name] || name
return modelNames[modelType] || modelType
}
//
const getImportanceColor = (value) => {
if (value >= 0.5) return '#67C23A' //
if (value >= 0.2) return '#E6A23C' //
return '#F56C6C' //
}
//
const modelMetrics = computed(() => {
//
const modelResults = computed(() => {
if (!trainingResult.value?.metrics) return []
return Object.entries(trainingResult.value.metrics).map(([model, metrics]) => ({
@ -257,15 +260,74 @@ const modelMetrics = computed(() => {
}))
})
//
const formatMoney = (value) => {
if (value === null || value === undefined) return '-'
return `${value.toFixed(2)} 元 (预测误差)`
//
const featureNameMap = {
//
'length_m': '总长(m)',
'width_m': '宽度(m)',
'height_m': '高度(m)',
'weight_kg': '重量(kg)',
'max_range_km': '最大射程(km)',
//
'firing_angle_horizontal': '方向射界(度)',
'firing_angle_vertical': '高低射界(度)',
'rocket_length_m': '火箭弹长度(m)',
'rocket_diameter_mm': '口径(mm)',
'rocket_weight_kg': '火箭弹重量(kg)',
'rate_of_fire': '射速(发/分)',
'combat_weight_kg': '战斗重量(kg)',
'speed_kmh': '速度(km/h)',
'min_range_km': '最小射程(km)',
'power_hp': '功率(hp)',
//
'fire_density': '火力密度',
'mobility_index': '机动性指标',
'range_ratio': '射程比',
'power_weight_ratio': '功重比',
'volume_density': '体积密度',
//
'wingspan_m': '翼展(m)',
'warhead_weight_kg': '战斗部重量(kg)',
'max_speed_ms': '最大速度(m/s)',
'cruise_speed_kmh': '巡航速度(km/h)',
'flight_time_min': '巡飞时间(min)',
'folded_length_mm': '折叠长度(mm)',
'folded_width_mm': '折叠宽度(mm)',
'folded_height_mm': '折叠高度(mm)',
//
'warhead_ratio': '战斗部比重',
'speed_ratio': '速度比',
'range_time_ratio': '射程时间比',
'aspect_ratio': '展弦比'
}
//
const featureImportanceData = computed(() => {
if (!trainingResult.value?.feature_importance || !trainingResult.value?.feature_names) return []
//
const data = trainingResult.value.feature_importance
.map((importance, index) => ({
feature: featureNameMap[trainingResult.value.feature_names[index]] || trainingResult.value.feature_names[index],
importance
}))
// 0
.filter(item => item.importance > 0)
//
.sort((a, b) => b.importance - a.importance)
return data
})
//
onMounted(() => {
//
if (trainingConfig.value.type) {
loadDatasets()
}
})
</script>
@ -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;
}
}
}
}
}
}
</style>

66
run.py
View File

@ -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
app.run(host='0.0.0.0', port=5001, debug=True)

View File

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

View File

@ -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'}
if __name__ == '__main__':
app = create_app()
app.run(host='localhost', port=5001)

View File

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

View File

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

View File

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

View File

@ -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()
if connection and connection.is_connected():
connection.close()
logger.info("Database connection closed")

View File

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

View File

@ -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)}")
logger.error(f"Import failed: {str(e)}")

33
src/logger.py Normal file
View File

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

View File

@ -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}")
# 标准化
# <EFBFBD><EFBFBD>标准化
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
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)
# 获取特<E58F96><E789B9><EFBFBD>重要性
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 {}

View File

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

View File

@ -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
# 预<><E9A284><EFBFBD>成本
# 预测成本
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<EFBFBD><EFBFBD><EFBFBD>为 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/<int:id>', 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/<int:id>', 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/<int:id>', 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/<int:id>', 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/<int:id>', 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/<int:id>', 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/<equipment_type>/latest', methods=['GET'])
def get_latest_model(equipment_type):
"""
获取最新训练的<EFBFBD><EFBFBD><EFBFBD>型信息
获取最新训练的型信息
"""
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/<int:id>/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/<int:id>', 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

View File

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