Update prediction results workflow
This commit is contained in:
+143
-116
@@ -61,6 +61,16 @@ PIPE_AGE_COLUMN = "管龄(年)"
|
||||
LEGACY_PIPE_AGE_COLUMN = "管龄"
|
||||
STATUS_COLUMN = "状态"
|
||||
EVENT_AGE_COLUMN = "事件/观察管龄(年)"
|
||||
COLUMN_ALIASES = {
|
||||
"管径(mm)": "管径",
|
||||
"管径(mm)": "管径",
|
||||
"流速(m/s)": "流速",
|
||||
"压力(MPa)": "压力",
|
||||
"压力(MPa)": "压力",
|
||||
"温度(℃)": "温度",
|
||||
"年均降雨量(mm)": "降雨量",
|
||||
"降雨量(mm)": "降雨量",
|
||||
}
|
||||
INPUT_COLUMNS = [
|
||||
ID_COLUMN,
|
||||
PIPE_AGE_COLUMN,
|
||||
@@ -68,6 +78,7 @@ INPUT_COLUMNS = [
|
||||
EVENT_AGE_COLUMN,
|
||||
*FEATURES,
|
||||
]
|
||||
CHART_DISPLAY_LIMIT = 10
|
||||
MATERIAL_COLUMN = "管材"
|
||||
MATERIAL_CODE_OPTIONS = [
|
||||
(1, "镀锌"),
|
||||
@@ -103,6 +114,7 @@ MATERIAL_ALIAS_TO_CODE.update(
|
||||
)
|
||||
SUPPORTED_EXTENSIONS = {".csv", ".xls", ".xlsx"}
|
||||
CHINESE_FONT_PROP = None
|
||||
DEFECT_GRADE_VALUES = {"无": 0.0, "轻度": 1.0, "中度": 3.0, "严重": 5.0}
|
||||
|
||||
|
||||
class PredictionError(Exception):
|
||||
@@ -119,12 +131,35 @@ class PredictionArtifacts:
|
||||
excel_path: Path
|
||||
image_path: Path
|
||||
image_filename: str
|
||||
importance_filename: str | None
|
||||
sample_count: int
|
||||
summary_rows: list[dict[str, Any]]
|
||||
analysis_text: str
|
||||
|
||||
|
||||
class ModelBundleAdapter:
|
||||
"""Expose a bundled preprocessor and survival model as one predictor."""
|
||||
|
||||
def __init__(self, preprocessor, model) -> None:
|
||||
self.preprocessor = preprocessor
|
||||
self.model = model
|
||||
|
||||
def _transform(self, frame: pd.DataFrame):
|
||||
values = frame[FEATURES].copy()
|
||||
values[MATERIAL_COLUMN] = values[MATERIAL_COLUMN].map(
|
||||
lambda value: f"M{int(value)}" if pd.notna(value) else value
|
||||
)
|
||||
values["位置"] = values["位置"].map(
|
||||
lambda value: f"L{int(value)}" if pd.notna(value) else value
|
||||
)
|
||||
return self.preprocessor.transform(values.to_numpy(dtype=object)).astype(np.float32)
|
||||
|
||||
def predict_survival_function(self, frame: pd.DataFrame):
|
||||
return self.model.predict_survival_function(self._transform(frame))
|
||||
|
||||
def predict(self, frame: pd.DataFrame):
|
||||
return self.model.predict(self._transform(frame))
|
||||
|
||||
|
||||
def configure_matplotlib_fonts():
|
||||
for font_path in CHINESE_FONT_FILES:
|
||||
path = Path(font_path)
|
||||
@@ -159,6 +194,10 @@ def load_model(model_path: str):
|
||||
if not os.path.exists(model_path):
|
||||
raise FileNotFoundError(f"未找到模型文件: {model_path}")
|
||||
model = joblib.load(model_path)
|
||||
if isinstance(model, dict):
|
||||
if not {"preprocessor", "model"}.issubset(model):
|
||||
raise ValueError("模型包缺少 preprocessor 或 model。")
|
||||
model = ModelBundleAdapter(model["preprocessor"], model["model"])
|
||||
feature_names = getattr(model, "feature_names_in_", None)
|
||||
if feature_names is not None and list(feature_names) != FEATURES:
|
||||
raise ValueError(
|
||||
@@ -195,17 +234,25 @@ def read_input_file(path: Path, suffix: str) -> pd.DataFrame:
|
||||
if suffix == ".csv":
|
||||
df = pd.read_csv(path, dtype={ID_COLUMN: "string"})
|
||||
else:
|
||||
df = pd.read_excel(path, dtype={ID_COLUMN: "string"})
|
||||
try:
|
||||
df = pd.read_excel(path, sheet_name="Template", dtype={ID_COLUMN: "string"})
|
||||
except ValueError:
|
||||
df = pd.read_excel(path, dtype={ID_COLUMN: "string"})
|
||||
df["_excel_row"] = df.index + 2
|
||||
except Exception as exc:
|
||||
logging.exception("文件解析失败: %s", exc)
|
||||
raise PredictionError("文件解析失败,请检查编码或表格格式。")
|
||||
return normalize_input_columns(df)
|
||||
df = normalize_input_columns(df)
|
||||
if suffix in {".xls", ".xlsx"}:
|
||||
df = fill_defects_from_detail_sheet(path, df)
|
||||
return df
|
||||
|
||||
|
||||
def normalize_input_columns(df: pd.DataFrame) -> pd.DataFrame:
|
||||
aliases = dict(COLUMN_ALIASES)
|
||||
if PIPE_AGE_COLUMN not in df.columns and LEGACY_PIPE_AGE_COLUMN in df.columns:
|
||||
return df.rename(columns={LEGACY_PIPE_AGE_COLUMN: PIPE_AGE_COLUMN})
|
||||
return df
|
||||
aliases[LEGACY_PIPE_AGE_COLUMN] = PIPE_AGE_COLUMN
|
||||
return df.rename(columns={name: aliases[name] for name in df.columns if name in aliases})
|
||||
|
||||
|
||||
def normalize_id_value(value: Any, fallback: str) -> str:
|
||||
@@ -216,6 +263,79 @@ def normalize_id_value(value: Any, fallback: str) -> str:
|
||||
return str(value)
|
||||
|
||||
|
||||
def defect_grade_value(value: Any) -> float | None:
|
||||
if pd.isna(value):
|
||||
return None
|
||||
text = str(value).strip()
|
||||
if text == "":
|
||||
return None
|
||||
if text in DEFECT_GRADE_VALUES:
|
||||
return DEFECT_GRADE_VALUES[text]
|
||||
try:
|
||||
return float(text)
|
||||
except ValueError:
|
||||
return None
|
||||
|
||||
|
||||
def weighted_defect_score(values: list[Any], weights: list[float]) -> float | None:
|
||||
scores = [defect_grade_value(value) for value in values]
|
||||
if any(score is None for score in scores):
|
||||
return None
|
||||
return round(sum(score * weight for score, weight in zip(scores, weights)), 3)
|
||||
|
||||
|
||||
def numeric_or_none(value: Any) -> float | None:
|
||||
try:
|
||||
if pd.isna(value):
|
||||
return None
|
||||
return float(value)
|
||||
except (TypeError, ValueError):
|
||||
return None
|
||||
|
||||
|
||||
def fill_defects_from_detail_sheet(path: Path, df: pd.DataFrame) -> pd.DataFrame:
|
||||
if "结构缺陷" not in df.columns or "功能缺陷" not in df.columns or "_excel_row" not in df.columns:
|
||||
return df
|
||||
|
||||
structure = pd.to_numeric(df["结构缺陷"], errors="coerce")
|
||||
function = pd.to_numeric(df["功能缺陷"], errors="coerce")
|
||||
needs_fill = structure.isna() | function.isna()
|
||||
df["结构缺陷"] = structure
|
||||
df["功能缺陷"] = function
|
||||
if not needs_fill.any():
|
||||
return df
|
||||
|
||||
try:
|
||||
from openpyxl import load_workbook
|
||||
|
||||
workbook = load_workbook(path, data_only=False, read_only=True)
|
||||
defect_sheet = workbook["缺陷计算"]
|
||||
except Exception as exc:
|
||||
logging.info("未能读取缺陷计算工作表,跳过缺陷自动计算: %s", exc)
|
||||
return df
|
||||
|
||||
for index in df.index[needs_fill]:
|
||||
excel_row = int(df.at[index, "_excel_row"])
|
||||
if pd.isna(df.at[index, "结构缺陷"]):
|
||||
value = numeric_or_none(defect_sheet.cell(excel_row, 5).value)
|
||||
if value is None:
|
||||
value = weighted_defect_score(
|
||||
[defect_sheet.cell(excel_row, column).value for column in (2, 3, 4)],
|
||||
[0.5, 0.4, 0.1],
|
||||
)
|
||||
df.at[index, "结构缺陷"] = value
|
||||
if pd.isna(df.at[index, "功能缺陷"]):
|
||||
value = numeric_or_none(defect_sheet.cell(excel_row, 10).value)
|
||||
if value is None:
|
||||
value = weighted_defect_score(
|
||||
[defect_sheet.cell(excel_row, column).value for column in (6, 7, 8, 9)],
|
||||
[0.5, 0.2, 0.2, 0.1],
|
||||
)
|
||||
df.at[index, "功能缺陷"] = value
|
||||
|
||||
return df
|
||||
|
||||
|
||||
def validate_input_frame(df: pd.DataFrame) -> None:
|
||||
if df.empty:
|
||||
raise PredictionError("上传文件没有可预测的数据。")
|
||||
@@ -288,99 +408,6 @@ def make_analysis_text(summary_rows: list[dict[str, Any]]) -> str:
|
||||
)
|
||||
|
||||
|
||||
def compute_feature_importance(model, x_test: pd.DataFrame) -> np.ndarray | None:
|
||||
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)
|
||||
vals = np.zeros(len(FEATURES), dtype=float)
|
||||
for j, feat in enumerate(FEATURES):
|
||||
diffs = []
|
||||
for _ in range(5):
|
||||
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
|
||||
|
||||
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)
|
||||
|
||||
fig, ax = plt.subplots(figsize=(9, max(3.0, 0.62 * n + 1.6)))
|
||||
cmap = LinearSegmentedColormap.from_list("brand", ["#7fb2e6", "#005EB8", "#0c4188"])
|
||||
colors = cmap(np.linspace(0.15, 1.0, n)) if n else None
|
||||
y_pos = np.arange(n)
|
||||
bars = ax.barh(y_pos, sorted_vals, color=colors, height=0.66, edgecolor="white", linewidth=0.8, zorder=3)
|
||||
|
||||
font_kwargs = chinese_font_kwargs()
|
||||
ax.set_xlabel("相对重要性", fontsize=11, color="#475569", **font_kwargs)
|
||||
ax.set_title("模型输入因素重要性排序", fontsize=15, fontweight="bold", color="#0f172a", pad=14, **font_kwargs)
|
||||
ax.set_yticks(y_pos)
|
||||
ax.set_yticklabels(sorted_feats, **font_kwargs)
|
||||
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)
|
||||
if CHINESE_FONT_PROP:
|
||||
for label in [*ax.get_yticklabels(), *ax.get_xticklabels()]:
|
||||
label.set_fontproperties(CHINESE_FONT_PROP)
|
||||
|
||||
max_val = float(sorted_vals.max()) if n else 0.0
|
||||
for bar, value in zip(bars, sorted_vals):
|
||||
ax.text(
|
||||
bar.get_width() + max_val * 0.012,
|
||||
bar.get_y() + bar.get_height() / 2,
|
||||
f"{value:.1%}",
|
||||
va="center",
|
||||
ha="left",
|
||||
fontsize=10,
|
||||
fontweight="bold",
|
||||
color="#1e293b",
|
||||
**font_kwargs,
|
||||
)
|
||||
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)
|
||||
|
||||
|
||||
def run_prediction(uploaded: FileStorage, user_id: int, model) -> PredictionArtifacts:
|
||||
original_filename = uploaded.filename or ""
|
||||
if not original_filename:
|
||||
@@ -409,15 +436,6 @@ def run_prediction(uploaded: FileStorage, user_id: int, model) -> PredictionArti
|
||||
image_path = IMAGE_DIR / image_filename
|
||||
summary_rows, summary_sheet_rows = render_survival_chart(df, curves, image_path)
|
||||
|
||||
importance_filename = None
|
||||
try:
|
||||
importance_values = compute_feature_importance(model, x_test)
|
||||
if importance_values is not None:
|
||||
importance_filename = f"importance_{user_id}_{run_id}.png"
|
||||
render_importance_chart(importance_values, IMAGE_DIR / importance_filename)
|
||||
except Exception as exc:
|
||||
logging.exception("生成特征重要性图失败: %s", exc)
|
||||
|
||||
safe_stem = Path(saved_filename).stem
|
||||
excel_path = user_dir / f"{safe_stem}_pre.xlsx"
|
||||
write_prediction_workbook(excel_path, curves, summary_rows, summary_sheet_rows)
|
||||
@@ -427,7 +445,6 @@ def run_prediction(uploaded: FileStorage, user_id: int, model) -> PredictionArti
|
||||
excel_path=excel_path,
|
||||
image_path=image_path,
|
||||
image_filename=image_filename,
|
||||
importance_filename=importance_filename,
|
||||
sample_count=len(summary_rows),
|
||||
summary_rows=summary_rows,
|
||||
analysis_text=make_analysis_text(summary_rows),
|
||||
@@ -441,13 +458,17 @@ def render_survival_chart(df: pd.DataFrame, curves, image_path: Path) -> tuple[l
|
||||
plt.figure(figsize=(10, 5.6))
|
||||
summary_rows: list[dict[str, Any]] = []
|
||||
summary_sheet_rows: list[dict[str, Any]] = []
|
||||
display_note = f"说明:图中仅展示前{CHART_DISPLAY_LIMIT}条管道的示例数据,不足{CHART_DISPLAY_LIMIT}条则全部展示;完整结果请下载电子表格。"
|
||||
|
||||
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 = normalize_id_value(df.iloc[i][ID_COLUMN], f"Pipe_{i+1:03d}")
|
||||
pipe_age = f"{df.iloc[i][PIPE_AGE_COLUMN]} 年" if PIPE_AGE_COLUMN in df.columns and pd.notna(df.iloc[i][PIPE_AGE_COLUMN]) else "-"
|
||||
health_probability = interpolate_probability(times, probs, 10)
|
||||
raw_pipe_age = df.iloc[i][PIPE_AGE_COLUMN] if PIPE_AGE_COLUMN in df.columns else None
|
||||
pipe_age_value = float(raw_pipe_age) if pd.notna(raw_pipe_age) else 0.0
|
||||
pipe_age = f"{raw_pipe_age} 年" if pd.notna(raw_pipe_age) else "-"
|
||||
health_probability = interpolate_probability(times, probs, pipe_age_value)
|
||||
health_risk = 1.0 - health_probability
|
||||
remaining_life = estimate_remaining_life(times, probs)
|
||||
grade_label, grade_desc, grade_class = grade_info(health_probability)
|
||||
|
||||
@@ -456,6 +477,7 @@ def render_survival_chart(df: pd.DataFrame, curves, image_path: Path) -> tuple[l
|
||||
"pipe_id": pipe_id,
|
||||
"pipe_age": pipe_age,
|
||||
"health_probability": health_probability,
|
||||
"health_risk": health_risk,
|
||||
"remaining_life": remaining_life,
|
||||
"grade_label": grade_label,
|
||||
"grade_desc": grade_desc,
|
||||
@@ -466,23 +488,25 @@ def render_survival_chart(df: pd.DataFrame, curves, image_path: Path) -> tuple[l
|
||||
{
|
||||
ID_COLUMN: pipe_id,
|
||||
PIPE_AGE_COLUMN: pipe_age,
|
||||
"健康等级": grade_label,
|
||||
"健康风险值": health_risk,
|
||||
}
|
||||
)
|
||||
plt.step(times, probs, where="post", linewidth=2, label=pipe_id)
|
||||
if i < CHART_DISPLAY_LIMIT:
|
||||
plt.step(times, probs, where="post", linewidth=2, label=pipe_id)
|
||||
|
||||
font_kwargs = chinese_font_kwargs()
|
||||
plt.xlabel("管龄(年)", **font_kwargs)
|
||||
plt.ylabel("健康风险", **font_kwargs)
|
||||
plt.title("预测分析图", **font_kwargs)
|
||||
plt.title("管道的剩余寿命分析图", **font_kwargs)
|
||||
plt.figtext(0.5, 0.02, display_note, ha="center", fontsize=9, color="#475569", **font_kwargs)
|
||||
plt.grid(alpha=0.18)
|
||||
if len(summary_rows) <= 12:
|
||||
if min(len(summary_rows), CHART_DISPLAY_LIMIT) <= 12:
|
||||
plt.legend(loc="best", fontsize=8, prop=CHINESE_FONT_PROP)
|
||||
ax = plt.gca()
|
||||
if CHINESE_FONT_PROP:
|
||||
for label in [*ax.get_xticklabels(), *ax.get_yticklabels()]:
|
||||
label.set_fontproperties(CHINESE_FONT_PROP)
|
||||
plt.tight_layout()
|
||||
plt.tight_layout(rect=(0, 0.07, 1, 1))
|
||||
plt.savefig(image_path, dpi=160, bbox_inches="tight")
|
||||
plt.close()
|
||||
return summary_rows, summary_sheet_rows
|
||||
@@ -511,7 +535,7 @@ def write_prediction_workbook(
|
||||
]
|
||||
|
||||
with pd.ExcelWriter(excel_path, engine="openpyxl") as writer:
|
||||
summary_df = pd.DataFrame(summary_sheet_rows, columns=[ID_COLUMN, PIPE_AGE_COLUMN, "健康等级"])
|
||||
summary_df = pd.DataFrame(summary_sheet_rows, columns=[ID_COLUMN, PIPE_AGE_COLUMN, "健康风险值"])
|
||||
sample_df = pd.DataFrame(sample_data_rows, columns=sample_columns)
|
||||
if ID_COLUMN in summary_df.columns:
|
||||
summary_df[ID_COLUMN] = summary_df[ID_COLUMN].astype("string")
|
||||
@@ -520,6 +544,9 @@ def write_prediction_workbook(
|
||||
sample_df.to_excel(writer, sheet_name="样本数据", index=False)
|
||||
|
||||
sample_worksheet = writer.book["样本数据"]
|
||||
summary_worksheet = writer.book["结果摘要"]
|
||||
for cell in summary_worksheet["C"][1:]:
|
||||
cell.number_format = "0.0%"
|
||||
sample_worksheet.freeze_panes = "A2"
|
||||
sample_worksheet.auto_filter.ref = sample_worksheet.dimensions
|
||||
sample_worksheet.column_dimensions["A"].width = 12
|
||||
|
||||
@@ -130,16 +130,10 @@ def password_reset_users():
|
||||
|
||||
def prediction_result_payload(artifacts, record: UploadRecord) -> dict:
|
||||
image_url = url_for("static", filename=f"images/{artifacts.image_filename}")
|
||||
importance_url = (
|
||||
url_for("static", filename=f"images/{artifacts.importance_filename}")
|
||||
if artifacts.importance_filename
|
||||
else None
|
||||
)
|
||||
return {
|
||||
"original_filename": artifacts.original_filename,
|
||||
"generated_at": format_app_datetime(record.upload_time),
|
||||
"image_url": image_url,
|
||||
"importance_url": importance_url,
|
||||
"excel_url": url_for("main.download_file", record_id=record.id, file_type="prediction"),
|
||||
"result_url": url_for("main.result_page"),
|
||||
"sample_count": int(artifacts.sample_count),
|
||||
@@ -439,7 +433,6 @@ def predict():
|
||||
{
|
||||
"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"],
|
||||
|
||||
Reference in New Issue
Block a user