# -*- coding: utf-8 -*-
"""
运行要求:
1. 同目录放置 my_survival_forest_model_quxi-10-0331.joblib
2. pip install flask flask_sqlalchemy flask_login pandas joblib matplotlib openpyxl xlsxwriter xlrd
3. python final_flask_app_0331_strict.py
"""
from __future__ import annotations
import logging
import os
import random
import sys
from dataclasses import dataclass
from datetime import datetime
from io import BytesIO
from pathlib import Path
from typing import Any, Dict, List
import joblib
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
import numpy as np
import pandas as pd
from flask import (
Flask,
abort,
flash,
jsonify,
redirect,
render_template_string,
request,
send_file,
session,
url_for,
)
from flask_login import (
LoginManager,
UserMixin,
current_user,
login_required,
login_user,
logout_user,
)
from flask_sqlalchemy import SQLAlchemy
from matplotlib import font_manager, rcParams
from werkzeug.security import check_password_hash, generate_password_hash
from werkzeug.utils import secure_filename
CHINESE_FONT_CANDIDATES = [
"Noto Sans CJK SC",
"Noto Sans CJK JP",
"Noto Sans CJK TC",
"Source Han Sans SC",
"WenQuanYi Micro Hei",
"SimHei",
"Microsoft YaHei",
"Arial Unicode MS",
]
def configure_matplotlib_fonts() -> None:
available_fonts = {font.name for font in font_manager.fontManager.ttflist}
selected_fonts = [font for font in CHINESE_FONT_CANDIDATES if font in available_fonts]
rcParams["font.sans-serif"] = selected_fonts + ["DejaVu Sans"]
configure_matplotlib_fonts()
rcParams["axes.unicode_minus"] = False
BASE_DIR = Path(__file__).resolve().parent
DATA_DIR = BASE_DIR / "data"
STATIC_DIR = BASE_DIR / "static"
UPLOAD_DIR = BASE_DIR / "uploads"
IMAGE_DIR = STATIC_DIR / "images"
FEATURES = [
"管材", "管径", "流速", "压力",
"温度", "降雨量", "位置",
"结构缺陷", "功能缺陷",
]
ID_COLUMN = "管道编号"
TEMPLATE_COLUMNS = [
ID_COLUMN,
"管龄", "状态", "管材", "管径", "流速",
"压力", "温度", "降雨量", "位置",
"结构缺陷", "功能缺陷",
]
DATA_DIR.mkdir(parents=True, exist_ok=True)
logging.basicConfig(
filename=str(DATA_DIR / "app.log"),
level=logging.INFO,
format="%(asctime)s - %(levelname)s - %(message)s",
)
app = Flask(__name__)
app.config["SECRET_KEY"] = "pipe-survival-0331-strict-secret"
app.config["SQLALCHEMY_DATABASE_URI"] = f"sqlite:///{DATA_DIR / 'pipe_survival_0331.db'}"
app.config["SQLALCHEMY_TRACK_MODIFICATIONS"] = False
app.config["UPLOAD_FOLDER"] = str(UPLOAD_DIR)
db = SQLAlchemy(app)
login_manager = LoginManager(app)
login_manager.login_view = "login"
login_manager.login_message = "请先登录后再访问该页面。"
class User(UserMixin, db.Model):
__tablename__ = "users"
id = db.Column(db.Integer, primary_key=True)
username = db.Column(db.String(100), unique=True, nullable=False)
password_hash = db.Column(db.String(255), nullable=False)
is_admin = db.Column(db.Boolean, default=False, nullable=False)
created_at = db.Column(db.DateTime, default=datetime.utcnow)
def set_password(self, password: str) -> None:
self.password_hash = generate_password_hash(password)
def check_password(self, password: str) -> bool:
return check_password_hash(self.password_hash, password)
class UploadRecord(db.Model):
__tablename__ = "upload_records"
id = db.Column(db.Integer, primary_key=True)
user_id = db.Column(db.Integer, db.ForeignKey("users.id"), nullable=False)
original_filename = db.Column(db.String(255), nullable=False)
saved_path = db.Column(db.String(500), nullable=False)
prediction_path = db.Column(db.String(500), nullable=False)
image_path = db.Column(db.String(500), nullable=False)
upload_time = db.Column(db.DateTime, default=datetime.utcnow)
user = db.relationship("User", backref=db.backref("uploads", lazy=True))
@login_manager.user_loader
def load_user(user_id: str):
return db.session.get(User, int(user_id))
@dataclass
class PipeSummary:
pipe_id: str
pipe_age: str
grade_label: str
def resource_path(relative_path: str) -> str:
try:
base_path = Path(sys._MEIPASS) # type: ignore[attr-defined]
except Exception:
base_path = BASE_DIR
return str((base_path / relative_path).resolve())
def ensure_dirs() -> None:
for path in [DATA_DIR, STATIC_DIR, IMAGE_DIR, UPLOAD_DIR]:
path.mkdir(parents=True, exist_ok=True)
def init_app() -> None:
ensure_dirs()
with app.app_context():
db.create_all()
if not User.query.filter_by(username="admin").first():
admin = User(username="admin", is_admin=True)
admin.set_password("admin123")
db.session.add(admin)
db.session.commit()
try:
app.config["RSF_MODEL"] = load_model()
logging.info("模型加载成功")
except Exception as exc:
app.config["RSF_MODEL"] = None
logging.exception("模型加载失败: %s", exc)
def load_model():
model_path = resource_path("my_survival_forest_model_quxi-10-0331.joblib")
if not os.path.exists(model_path):
raise FileNotFoundError(f"未找到模型文件: {model_path}")
return joblib.load(model_path)
@app.context_processor
def inject_helpers():
def now_year() -> int:
return datetime.now().year
return {"feature_list": FEATURES, "now_year": now_year}
BASE_TEMPLATE_HEAD = r"""
"""
LOGIN_TEMPLATE = r"""
{{ '注册' if mode == 'register' else '登录' }} |
""" + BASE_TEMPLATE_HEAD + r"""
water_drop
供水管道健康评估系统
上传管网数据,自动生成健康状态评估、剩余寿命预测与输入因素重要性分析。
© {{ now_year() }} 供水管道健康评估系统
{% with messages = get_flashed_messages(with_categories=true) %}
{% if messages %}
{% for category, message in messages %}
{{ message }}
{% endfor %}
{% endif %}
{% endwith %}
{% if mode == 'login' %}
{% else %}
{% endif %}
"""
DASHBOARD_TEMPLATE = r"""
""" + BASE_TEMPLATE_HEAD + r"""
notifications
help
{{ current_user.username[:1]|upper }}
退出登录
{% with messages = get_flashed_messages(with_categories=true) %}
{% if messages %}
{% for category, message in messages %}
{{ message }}
{% endfor %}
{% endif %}
{% endwith %}
生存概率阶梯图
模型输入因素重要性排序
结果摘要
上传文件
预测样本数-
模型my_survival_forest_model_quxi-10-0331
"""
RESULT_TEMPLATE = r"""
预测结果
""" + BASE_TEMPLATE_HEAD + r"""
notifications
help
{{ current_user.username[:1]|upper }}
退出登录
{% if result %}
description
数据分析报告: {{ result.original_filename }}
基于上传数据生成的实时分析报告 • 生成于: {{ result.generated_at }}
(0, 0.2]
管道安全风险十分严重,需立刻进行抢修或更新改造
(0.2, 0.4]
管道安全风险较为严重,需尽快安排检修及加频巡检
(0.4, 0.6]
管道安全风险较低,需安排定期巡检
(0.6, 0.8]
管道安全风险较小,维持常规巡视
| 管道编号 |
健康等级 |
{% for item in result.summary_rows %}
| {{ item.pipe_id }} |
{{ item.grade_label }}
|
{% endfor %}
共 {{ result.sample_count }} 个样本,显示 {{ result.summary_rows|length }} 条
查看完整样本列表
info
分析说明:{{ result.analysis_text }}
{% if result.importance_url %}
insights
说明:柱状图按重要性从高到低展示各输入因素对模型预测结果的相对贡献,数值为归一化占比,可用于辅助识别影响管道健康状态的关键因素。
{% endif %}
{% else %}
当前还没有预测结果,请先从主页上传文件并运行预测。
{% endif %}
"""
ADMIN_TEMPLATE = r"""
管理台
""" + BASE_TEMPLATE_HEAD + r"""
| 用户 |
原始文件 |
上传时间 |
下载 |
{% for record in records %}
| {{ record.user.username }} |
{{ record.original_filename }} |
{{ record.upload_time.strftime('%Y-%m-%d %H:%M:%S') }} |
原始文件
预测结果
|
{% else %}
| 暂无上传记录 |
{% endfor %}
"""
HISTORY_TEMPLATE = r"""
预测历史
""" + BASE_TEMPLATE_HEAD + r"""
{% for record in records %}
{{ record.original_filename }}
{{ record.upload_time.strftime('%Y-%m-%d %H:%M:%S') }}
{% else %}
还没有历史记录。
{% endfor %}
"""
def grade_info(probability: float) -> tuple[str, str, str]:
if probability <= 0.2:
return ("I级", "管道安全风险十分严重,需立刻进行抢修或更新改造", "bg-dangerSoft text-dangerText")
if probability <= 0.4:
return ("II级", "管道安全风险较为严重,需尽快安排检修及加频巡检", "bg-orange-50 text-orange-600")
if probability <= 0.6:
return ("III级", "管道安全风险较低,需安排定期巡检", "bg-amber-50 text-amber-600")
if probability <= 0.8:
return ("IV级", "管道安全风险较小,维持常规巡视", "bg-blue-50 text-blue-600")
return ("V级", "管道安全,维持常规巡视", "bg-blueSoft text-primary")
def interpolate_probability(times: List[float], probs: List[float], target: float) -> float:
if not times:
return 0.0
if target <= times[0]:
return float(probs[0])
for idx in range(1, len(times)):
if times[idx] >= target:
return float(probs[idx])
return float(probs[-1])
def estimate_remaining_life(times: List[float], probs: List[float]) -> float:
for t, p in zip(times, probs):
if p <= 0.5:
return float(t)
return float(times[-1]) if times else 0.0
def make_analysis_text(summary_rows: List[Dict[str, Any]]) -> str:
if not summary_rows:
return "当前结果为空,暂无可供解释的样本。"
worst = min(summary_rows, key=lambda x: x["health_probability"])
best = max(summary_rows, key=lambda x: x["health_probability"])
return (
f"阶梯状曲线表示模型对不同管道随时间推移维持在安全健康状态概率的动态预测。"
)
def compute_feature_importance(model, x_test: "pd.DataFrame") -> "np.ndarray | None":
"""返回各输入因素的相对重要性(与 FEATURES 顺序一致,归一化为占比)。
优先使用模型自带的 feature_importances_;若不可用,则采用与标签无关的
置换重要性:打乱单个特征后观察模型风险评分的平均变化幅度。
"""
try:
importances = model.feature_importances_
except Exception:
importances = None
if importances is not None and len(importances) == len(FEATURES):
vals = np.asarray(importances, dtype=float)
else:
try:
baseline = np.asarray(model.predict(x_test), dtype=float)
except Exception as exc:
logging.exception("特征重要性基线预测失败: %s", exc)
return None
rng = np.random.default_rng(42)
n_repeats = 5
vals = np.zeros(len(FEATURES), dtype=float)
for j, feat in enumerate(FEATURES):
diffs = []
for _ in range(n_repeats):
x_perm = x_test.copy()
x_perm[feat] = rng.permutation(x_perm[feat].to_numpy())
try:
perm_pred = np.asarray(model.predict(x_perm), dtype=float)
except Exception:
perm_pred = baseline
diffs.append(float(np.mean(np.abs(perm_pred - baseline))))
vals[j] = float(np.mean(diffs)) if diffs else 0.0
vals = np.clip(vals, a_min=0.0, a_max=None)
total = float(vals.sum())
if total > 0:
vals = vals / total
return vals
def render_importance_chart(values: "np.ndarray", save_path: "Path") -> None:
from matplotlib.colors import LinearSegmentedColormap
# 仅保留有实际贡献的因素(占比 >= 0.05%,即不会显示为 0.0%);若全部过低则回退展示全部
keep = values >= 5e-4
if not bool(np.any(keep)):
keep = np.ones_like(values, dtype=bool)
kept_feats = [FEATURES[i] for i in range(len(FEATURES)) if keep[i]]
kept_vals = values[keep]
order = np.argsort(kept_vals)
sorted_feats = [kept_feats[k] for k in order]
sorted_vals = kept_vals[order]
n = len(sorted_vals)
height = max(3.0, 0.62 * n + 1.6)
fig, ax = plt.subplots(figsize=(9, height))
cmap = LinearSegmentedColormap.from_list("brand", ["#7fb2e6", "#005EB8", "#0c4188"])
colors = cmap(np.linspace(0.15, 1.0, n)) if n else None
bars = ax.barh(sorted_feats, sorted_vals, color=colors, height=0.66,
edgecolor="white", linewidth=0.8, zorder=3)
ax.set_xlabel("相对重要性", fontsize=11, color="#475569")
ax.set_title("模型输入因素重要性排序", fontsize=15, fontweight="bold",
color="#0f172a", pad=14)
ax.grid(axis="x", color="#e2e8f0", linewidth=1, zorder=0)
ax.set_axisbelow(True)
for spine in ("top", "right", "left"):
ax.spines[spine].set_visible(False)
ax.spines["bottom"].set_color("#cbd5e1")
ax.tick_params(axis="y", length=0, labelsize=11)
ax.tick_params(axis="x", colors="#94a3b8", labelsize=9)
max_val = float(sorted_vals.max()) if n else 0.0
for bar, v in zip(bars, sorted_vals):
ax.text(bar.get_width() + max_val * 0.012,
bar.get_y() + bar.get_height() / 2,
f"{v:.1%}", va="center", ha="left",
fontsize=10, fontweight="bold", color="#1e293b")
if max_val > 0:
ax.set_xlim(0, max_val * 1.18)
fig.tight_layout()
fig.savefig(save_path, dpi=160, bbox_inches="tight")
plt.close(fig)
@app.route("/")
def index():
if current_user.is_authenticated:
return redirect(url_for("home"))
return redirect(url_for("login"))
@app.route("/login", methods=["GET", "POST"])
def login():
if current_user.is_authenticated:
return redirect(url_for("home"))
if request.method == "GET":
session["captcha"] = "".join(random.choices("ABCDEFGHJKLMNPQRSTUVWXYZ23456789", k=5))
return render_template_string(LOGIN_TEMPLATE, mode="login", captcha=session["captcha"])
username = request.form.get("username", "").strip()
password = request.form.get("password", "")
captcha_input = request.form.get("captcha", "").strip().upper()
if captcha_input != session.get("captcha", ""):
flash("验证码错误", "error")
session["captcha"] = "".join(random.choices("ABCDEFGHJKLMNPQRSTUVWXYZ23456789", k=5))
return render_template_string(LOGIN_TEMPLATE, mode="login", captcha=session["captcha"]), 400
user = User.query.filter_by(username=username).first()
if not user or not user.check_password(password):
flash("用户名或密码错误", "error")
session["captcha"] = "".join(random.choices("ABCDEFGHJKLMNPQRSTUVWXYZ23456789", k=5))
return render_template_string(LOGIN_TEMPLATE, mode="login", captcha=session["captcha"]), 400
login_user(user, remember=bool(request.form.get("remember")))
return redirect(url_for("home"))
@app.route("/register", methods=["GET", "POST"])
def register():
if request.method == "GET":
return render_template_string(LOGIN_TEMPLATE, mode="register", captcha="")
username = request.form.get("username", "").strip()
password = request.form.get("password", "")
if not username:
flash("用户名不能为空", "error")
return render_template_string(LOGIN_TEMPLATE, mode="register", captcha=""), 400
if len(password) < 6:
flash("密码至少需要 6 位", "error")
return render_template_string(LOGIN_TEMPLATE, mode="register", captcha=""), 400
if User.query.filter_by(username=username).first():
flash("用户名已存在", "error")
return render_template_string(LOGIN_TEMPLATE, mode="register", captcha=""), 400
user = User(username=username, is_admin=False)
user.set_password(password)
db.session.add(user)
db.session.commit()
flash("注册成功,请登录", "info")
session["captcha"] = "".join(random.choices("ABCDEFGHJKLMNPQRSTUVWXYZ23456789", k=5))
return render_template_string(LOGIN_TEMPLATE, mode="login", captcha=session["captcha"])
@app.route("/logout")
@login_required
def logout():
logout_user()
return redirect(url_for("login"))
@app.route("/home")
@login_required
def home():
return render_template_string(DASHBOARD_TEMPLATE)
@app.route("/history")
@login_required
def history_page():
records = UploadRecord.query.filter_by(user_id=current_user.id).order_by(UploadRecord.upload_time.desc()).all()
return render_template_string(HISTORY_TEMPLATE, records=records)
@app.route("/admin")
@login_required
def admin_dashboard():
if not current_user.is_admin:
abort(403)
records = UploadRecord.query.order_by(UploadRecord.upload_time.desc()).all()
return render_template_string(ADMIN_TEMPLATE, records=records)
@app.route("/download//")
@login_required
def download_file(record_id: int, file_type: str):
record = UploadRecord.query.get_or_404(record_id)
if not (current_user.is_admin or current_user.id == record.user_id):
abort(403)
file_path = record.saved_path if file_type == "original" else record.prediction_path
if not os.path.exists(file_path):
abort(404)
return send_file(file_path, as_attachment=True)
@app.route("/download_template")
def download_template():
template_path = BASE_DIR / "example.xlsx"
if not template_path.exists():
abort(404)
return send_file(template_path, as_attachment=True, download_name="example.xlsx")
@app.route("/result")
@login_required
def result_page():
result = session.get("last_result")
return render_template_string(RESULT_TEMPLATE, result=result)
@app.route("/predict", methods=["POST"])
@login_required
def predict():
model = app.config.get("RSF_MODEL")
if model is None:
return jsonify({"error": "模型未成功加载,请检查 my_survival_forest_model_quxi-10-0331.joblib 文件。"}), 500
uploaded = request.files.get("file")
if uploaded is None or uploaded.filename == "":
return jsonify({"error": "未选择文件"}), 400
filename = secure_filename(uploaded.filename)
if not filename.lower().endswith((".csv", ".xls", ".xlsx")):
return jsonify({"error": "不支持的格式,仅支持 CSV / XLS / XLSX"}), 400
user_dir = UPLOAD_DIR / f"user_{current_user.id}"
user_dir.mkdir(parents=True, exist_ok=True)
timestamp = datetime.now().strftime("%Y%m%d%H%M%S")
stem, ext = os.path.splitext(filename)
original_path = user_dir / f"{stem}_{timestamp}{ext}"
uploaded.save(original_path)
try:
if ext.lower() == ".csv":
df = pd.read_csv(original_path)
else:
df = pd.read_excel(original_path)
except Exception as exc:
logging.exception("文件解析失败: %s", exc)
return jsonify({"error": "文件解析失败,请检查编码或表格格式。"}), 400
missing = [col for col in FEATURES if col not in df.columns]
if missing:
return jsonify({"error": f"缺少必要字段: {', '.join(missing)}"}), 400
x_test = df[FEATURES].copy()
try:
curves = model.predict_survival_function(x_test)
except Exception as exc:
logging.exception("预测失败: %s", exc)
return jsonify({"error": "模型预测失败,请检查输入字段类型是否正确。"}), 500
plt.figure(figsize=(10, 5.6))
summary_rows: List[Dict[str, Any]] = []
summary_sheet_rows: List[Dict[str, Any]] = []
for i, curve in enumerate(curves):
times = [float(x) for x in list(curve.x)]
probs = [float(y) for y in list(curve.y)]
pipe_id = str(df.iloc[i]["Status"]) if "Status" in df.columns and pd.notna(df.iloc[i]["Status"]) else f"Pipe_{i+1:03d}"
pipe_age = f"{df.iloc[i]['Pipeage']} 年" if "Pipeage" in df.columns and pd.notna(df.iloc[i]["Pipeage"]) else "-"
health_probability = interpolate_probability(times, probs, 10)
remaining_life = estimate_remaining_life(times, probs)
grade_label, grade_desc, grade_class = grade_info(health_probability)
summary_rows.append({
"pipe_id": pipe_id,
"pipe_age": pipe_age,
"health_probability": health_probability,
"remaining_life": remaining_life,
"grade_label": grade_label,
"grade_desc": grade_desc,
"grade_class": grade_class,
})
summary_sheet_rows.append({
"管道标识": i + 1,
"健康等级": grade_label,
})
plt.step(times, probs, where="post", linewidth=2, label=pipe_id)
plt.xlabel("预测时间轴(年)")
plt.ylabel("生存概率")
plt.title("预测分析图")
plt.grid(alpha=0.18)
if len(summary_rows) <= 12:
plt.legend(loc="best", fontsize=8)
plt.tight_layout()
image_filename = f"plot_{current_user.id}_{timestamp}.png"
image_path = IMAGE_DIR / image_filename
plt.savefig(image_path, dpi=160, bbox_inches="tight")
plt.close()
importance_filename = None
try:
importance_values = compute_feature_importance(model, x_test)
if importance_values is not None:
importance_filename = f"importance_{current_user.id}_{timestamp}.png"
render_importance_chart(importance_values, IMAGE_DIR / importance_filename)
except Exception as exc:
logging.exception("生成特征重要性图失败: %s", exc)
importance_filename = None
excel_filename = f"{stem}_pre_{timestamp}.xlsx"
excel_path = user_dir / excel_filename
with pd.ExcelWriter(excel_path, engine="xlsxwriter") as writer:
pd.DataFrame(summary_sheet_rows).to_excel(writer, sheet_name="结果摘要", index=False)
for i, curve in enumerate(curves):
times = [float(x) for x in list(curve.x)]
probs = [float(y) for y in list(curve.y)]
pipe_id = summary_rows[i]["pipe_id"]
out_df = pd.DataFrame({"时间(年)": times, f"{pipe_id}生存概率": probs})
out_df.to_excel(writer, sheet_name=f"样本{i+1}", index=False)
record = UploadRecord(
user_id=current_user.id,
original_filename=filename,
saved_path=str(original_path),
prediction_path=str(excel_path),
image_path=str(image_path),
)
db.session.add(record)
db.session.commit()
analysis_text = make_analysis_text(summary_rows)
last_result = {
"original_filename": filename,
"generated_at": datetime.now().strftime("%Y-%m-%d %H:%M:%S"),
"image_url": url_for("static", filename=f"images/{image_filename}"),
"importance_url": url_for("static", filename=f"images/{importance_filename}") if importance_filename else None,
"excel_url": url_for("download_file", record_id=record.id, file_type="prediction"),
"result_url": url_for("result_page"),
"sample_count": int(len(summary_rows)),
"summary_rows": summary_rows[:3],
"analysis_text": analysis_text,
}
session["last_result"] = last_result
return jsonify({
"message": "预测成功",
"image_url": last_result["image_url"],
"importance_url": last_result["importance_url"],
"excel_url": last_result["excel_url"],
"result_url": last_result["result_url"],
"sample_count": last_result["sample_count"],
"original_filename": filename,
})
if __name__ == "__main__":
init_app()
app.run(host="0.0.0.0", port=5005, debug=False, use_reloader=False)