Skip to content

Physics-Informed Digital Twin Framework for RO Membrane Degradation Diagnosis #542

Description

@liujiaoying

import os
from pathlib import Path

import pandas as pd
import numpy as np
import torch
import torch.nn as nn
import matplotlib.pyplot as plt

=========================

0) 配置:数据路径 & sheet

=========================

PATH_A1 = "/Users/caixiaoliang/Downloads/d5ew00634a1.xlsx"
PATH_A2 = "/Users/caixiaoliang/Downloads/d5ew00634a2.xlsx"

SHEET_A1 = "膜污染A和B变化数据"
SHEET_A2 = "Subset of 10000 experimental da"

输出目录:默认输出到“当前脚本同目录 / outputs”

BASE_DIR = Path(file).resolve().parent if "file" in globals() else Path.cwd()
OUT_DIR = BASE_DIR / "outputs"
OUT_DIR.mkdir(parents=True, exist_ok=True)

SEED = 42
np.random.seed(SEED)
torch.manual_seed(SEED)

=========================

1) Time -> 秒(关键修复:先解析、再用 min 归零、并按时间说排序)

=========================

STRICT_TIME_FORMAT = False

def time_to_seconds(col: pd.Series) -> np.ndarray:
# 优先 HH:MM:SS
s = pd.to_datetime(col.astype(str), format="%H:%M:%S", errors="coerce")

# 兜底:让 pandas 自己推断(会有 warning,但不影响)
if s.isna().any():
    if STRICT_TIME_FORMAT:
        bad = col[s.isna()].head(5).tolist()
        raise ValueError(f"Time 存在非 HH:MM:SS 格式样本,例如:{bad}")
    s2 = pd.to_datetime(col.astype(str), errors="coerce")
    s = s.fillna(s2)

if s.isna().all():
    raise ValueError("Time 列无法解析为时间,请检查 Time 的格式。")

sec = s.dt.hour * 3600 + s.dt.minute * 60 + s.dt.second
sec = sec.astype(float)

# ✅关键:用 min 归零(防止首行不是最早时间导致出现负值)
sec = sec - sec.min()
return sec.to_numpy().reshape(-1, 1).astype(np.float32)

=========================

2) 渗透压模型(把你图里的公式真正放进来)

你给的经验式:π = 0.74 * TDS

注意:TDS 通常以 g/L 更合理;你的表是 mg/L,所以要 /1000

π 单位常见为 bar(工程经验式),而你的压力是 MPa,所以再 *0.1 转 MPa

=========================

def pi_from_tds_mgL(tds_mgL: torch.Tensor) -> torch.Tensor:
# mg/L -> g/L
tds_gL = tds_mgL / 1000.0
# π(bar) = 0.74 * TDS(g/L)
pi_bar = 0.74 * tds_gL
# bar -> MPa (1 bar = 0.1 MPa)
pi_mpa = 0.1 * pi_bar
return pi_mpa

=========================

3) 读数据 + 组装 X / obs

=========================

def build_xy_a1(df: pd.DataFrame):
df = df.copy()
t = time_to_seconds(df["Time"])
df["_t_sec"] = t.reshape(-1)

# ✅按时间排序,避免后续 t 乱序导致画图/保存异常
df = df.sort_values("_t_sec").reset_index(drop=True)
t = df["_t_sec"].to_numpy().reshape(-1, 1).astype(np.float32)

Pf = df["Feed Pressure(Mpa)"].to_numpy().reshape(-1, 1).astype(np.float32)
Pb = df["Brine Pressure(Mpa)"].to_numpy().reshape(-1, 1).astype(np.float32)
TDSf = df["Feed Concentration(mg/L)"].to_numpy().reshape(-1, 1).astype(np.float32)
T = df["Temperature(℃)"].to_numpy().reshape(-1, 1).astype(np.float32)

TDSp_obs = df["Permeate Concentration(mg/L)"].to_numpy().reshape(-1, 1).astype(np.float32)
Qp_obs = df["Permeate Flow Rate(L/h)"].to_numpy().reshape(-1, 1).astype(np.float32)

X = np.hstack([t, Pf, Pb, TDSf, T]).astype(np.float32)
obs = {
    "t_sec": t,
    "TDSp_obs": TDSp_obs,
    "Jw_obs": Qp_obs,
    "Pf": Pf,
    "Pb": Pb,
    "TDSf": TDSf
}
return X, obs

def build_xy_a2_primary(df: pd.DataFrame):
df = df.copy()

# Primary 稳态筛选(如果有该列就过滤)
switch_candidates = [
    "Primary switch", "Primary Switch", "Primary_Switch", "PrimarySwitch",
    "Primary mode", "Primary Mode", "Primary_Mode", "PrimaryMode"
]
used_switch = None
for c in switch_candidates:
    if c in df.columns:
        used_switch = c
        df = df[df[c] == 1]
        break

if used_switch:
    print(f"[a2] Applied stable-run filter using column: {used_switch}. Remaining rows: {len(df)}")
else:
    print("[a2] No Primary switch/mode column found; using all rows (may contain regime changes).")

t = time_to_seconds(df["Time"])
df["_t_sec"] = t.reshape(-1)
df = df.sort_values("_t_sec").reset_index(drop=True)
t = df["_t_sec"].to_numpy().reshape(-1, 1).astype(np.float32)

Pf = df["Primary Feed Pressure(Mpa)"].to_numpy().reshape(-1, 1).astype(np.float32)
Pb = df["Primary Brine Pressure(Mpa)"].to_numpy().reshape(-1, 1).astype(np.float32)
TDSf = df["Primary Feed Concentration(mg/L)"].to_numpy().reshape(-1, 1).astype(np.float32)
T = df["Temperature (℃)"].to_numpy().reshape(-1, 1).astype(np.float32)

TDSp_obs = df["Primary Permeate Concentration(mg/L)"].to_numpy().reshape(-1, 1).astype(np.float32)
Qp_obs = df["Primary Permeate Flow Rate(L/h)"].to_numpy().reshape(-1, 1).astype(np.float32)

X = np.hstack([t, Pf, Pb, TDSf, T]).astype(np.float32)
obs = {
    "t_sec": t,
    "TDSp_obs": TDSp_obs,
    "Jw_obs": Qp_obs,
    "Pf": Pf,
    "Pb": Pb,
    "TDSf": TDSf
}
return X, obs

=========================

4) 标准化(只对 X 做;物理计算用 obs 原始量)

=========================

def standardize(X: np.ndarray):
if np.isnan(X).any():
raise ValueError("X 中存在 NaN,请检查数据列是否有缺失值。")
mu = X.mean(axis=0, keepdims=True)
sd = X.std(axis=0, keepdims=True) + 1e-8
return (X - mu) / sd, mu, sd

=========================

5) PINN 网络:输出 [Jw, Js, A, B]

✅关键修复:A、B 必须为正(物理意义:透水系数/盐透过系数不能为负)

用 softplus 保证正值

=========================

class PINNNet(nn.Module):
def init(self, in_dim=5, hidden=64):
super().init()
self.backbone = nn.Sequential(
nn.Linear(in_dim, hidden), nn.Tanh(),
nn.Linear(hidden, hidden), nn.Tanh(),
nn.Linear(hidden, hidden), nn.Tanh(),
)
self.head = nn.Linear(hidden, 4)
self.sp = nn.Softplus(beta=1.0)

def forward(self, x):
    y = self.head(self.backbone(x))
    Jw_raw, Js_raw, A_raw, B_raw = y[:, 0:1], y[:, 1:2], y[:, 2:3], y[:, 3:4]

    # ✅保证正(也可只保证 A/B 正;这里 Jw/Js 也保证正,减少物理冲突)
    Jw = self.sp(Jw_raw)
    Js = self.sp(Js_raw)
    A = self.sp(A_raw)
    B = self.sp(B_raw)
    return Jw, Js, A, B

=========================

6) Loss:把你图里那套物理约束真正用上

(1) 数据项:Jw 贴合观测(你表里是 Permeate Flow Rate)

(2) 水通量物理:Jw = A * [(Pf - 1/2ΔP) - (πf-πp)]

其中 Pavg = Pf - 0.5*(Pf-Pb) = 0.5*(Pf+Pb)

(3) 盐通量物理:Js = B*(TDSf - TDSp)

(4) 平衡关系:TDSp = Js / Jw

✅这里我们用 “TDSp_pred = Js/Jw”,并让它贴合 TDSp_obs

=========================

def loss_terms(model, xn: torch.Tensor, obs: dict, lambdas=(1.0, 1.0, 1.0, 1.0), eps=1e-8):
lam_data, lam_flux, lam_salt, lam_tds = lambdas

Jw, Js, A, B = model(xn)

Pf = torch.tensor(obs["Pf"], dtype=torch.float32)
Pb = torch.tensor(obs["Pb"], dtype=torch.float32)
TDSf = torch.tensor(obs["TDSf"], dtype=torch.float32)
TDSp_obs = torch.tensor(obs["TDSp_obs"], dtype=torch.float32)
Jw_obs = torch.tensor(obs["Jw_obs"], dtype=torch.float32)

dP = Pf - Pb
Pavg = Pf - 0.5 * dP  # = 0.5*(Pf+Pb)

pi_f = pi_from_tds_mgL(TDSf)
pi_p = pi_from_tds_mgL(TDSp_obs)
dpi = pi_f - pi_p

# 物理:水通量
Jw_phy = A * (Pavg - dpi)

# 物理:盐通量
Js_phy = B * (TDSf - TDSp_obs)

# 平衡:TDSp_pred = Js/Jw
TDSp_pred = Js / (Jw + eps)

# 数据项(用相对误差更稳)
L_data = torch.mean(((Jw - Jw_obs) / (torch.abs(Jw_obs) + 1e-6)) ** 2)

# 物理残差
L_flux = torch.mean((Jw - Jw_phy) ** 2)
L_salt = torch.mean((Js - Js_phy) ** 2)

# 让 TDSp_pred 贴合观测 TDSp_obs(非常关键)
L_tds = torch.mean(((TDSp_pred - TDSp_obs) / (torch.abs(TDSp_obs) + 1e-6)) ** 2)

L_total = lam_data * L_data + lam_flux * L_flux + lam_salt * L_salt + lam_tds * L_tds
return L_total, (L_data.detach(), L_flux.detach(), L_salt.detach(), L_tds.detach())

@torch.no_grad()
def metrics_jw(model, xn: torch.Tensor, obs: dict, eps=1e-6):
Jw_pred, _, _, _ = model(xn)
y = torch.tensor(obs["Jw_obs"], dtype=torch.float32)
err = (Jw_pred - y).cpu().numpy().reshape(-1)
yy = y.cpu().numpy().reshape(-1)
mae = float(np.mean(np.abs(err)))
rmse = float(np.sqrt(np.mean(err ** 2)))
mape = float(np.mean(np.abs(err) / (np.abs(yy) + eps)) * 100.0)
return mae, rmse, mape

@torch.no_grad()
def predict_ab(model, xn: torch.Tensor):
Jw, Js, A, B = model(xn)
return (Jw.cpu().numpy(), Js.cpu().numpy(), A.cpu().numpy(), B.cpu().numpy())

=========================

7) 离线训练(a1)

=========================

def train_offline(model, Xn: np.ndarray, obs: dict, epochs=2000, lr=1e-3, lambdas=(1, 1, 1, 1)):
opt = torch.optim.Adam(model.parameters(), lr=lr)
x = torch.tensor(Xn, dtype=torch.float32)

for ep in range(1, epochs + 1):
    opt.zero_grad()
    L, (Ld, Lf, Ls, Lt) = loss_terms(model, x, obs, lambdas=lambdas)
    L.backward()
    opt.step()

    if ep % 200 == 0:
        mae, rmse, mape = metrics_jw(model, x, obs)
        print(f"[offline] ep={ep}  L={L.item():.4e}  "
              f"Ld={Ld.item():.4e} Lf={Lf.item():.4e} Ls={Ls.item():.4e} Lt={Lt.item():.4e}  "
              f"MAE={mae:.3f} RMSE={rmse:.3f} MAPE={mape:.2f}%")

=========================

8) 在线同步(a2)+ 导出 A(t)/B(t) & 图

=========================

def sync_online(
model,
Xn2: np.ndarray,
obs2: dict,
window=200,
stride=50,
steps=200,
lr=5e-5,
lambdas=(1, 1, 1, 1),
clip_norm=1.0,

# jump gating(窗口内部 max-min)
pf_jump_th=0.30,
tds_jump_th=2000.0,
jw_jump_th=80.0,

# distribution gating(z-score)
baseline_n=1000,
z_th_pf=6.0,
z_th_tds=6.0,
z_th_jw=6.0,

# sd floor(防止 sd 太小导致 z 爆表)
sd_floor_pf=0.02,
sd_floor_tds=30.0,
sd_floor_jw=5.0,

# 自动分段
reset_patience=6,

# 输出文件
out_ab_csv="AB_timeseries.csv",
out_regime_csv="regime_summary.csv",
out_plot_A="plot_A_t.png",
out_plot_B="plot_B_t.png",

):
opt2 = torch.optim.Adam(model.parameters(), lr=lr)
n = Xn2.shape[0]

# ----- 初始基线(a2 前 baseline_n)-----
bn = min(baseline_n, len(obs2["Pf"]))
mu_pf = float(obs2["Pf"][:bn].mean()); sd_pf = float(obs2["Pf"][:bn].std() + 1e-8)
mu_tds = float(obs2["TDSf"][:bn].mean()); sd_tds = float(obs2["TDSf"][:bn].std() + 1e-8)
mu_jw = float(obs2["Jw_obs"][:bn].mean()); sd_jw = float(obs2["Jw_obs"][:bn].std() + 1e-8)

sd_pf_eff = max(sd_pf, sd_floor_pf)
sd_tds_eff = max(sd_tds, sd_floor_tds)
sd_jw_eff = max(sd_jw, sd_floor_jw)

print(f"[baseline-a2:init] n={bn}  mu_pf={mu_pf:.4f} sd_pf={sd_pf:.4f}(eff={sd_pf_eff:.4f})  "
      f"mu_tds={mu_tds:.2f} sd_tds={sd_tds:.2f}(eff={sd_tds_eff:.2f})  "
      f"mu_jw={mu_jw:.2f} sd_jw={sd_jw:.2f}(eff={sd_jw_eff:.2f})")

ts_records = []
regime_records = []

consec_dist_fail = 0
regime_id = 0

reg_mae_list, reg_rmse_list, reg_mape_list, reg_loss_list = [], [], [], []

def flush_regime_summary(rid: int):
    if len(reg_loss_list) == 0:
        return
    regime_records.append({
        "regime_id": rid,
        "windows": len(reg_loss_list),
        "loss_mean": float(np.mean(reg_loss_list)),
        "loss_min": float(np.min(reg_loss_list)),
        "loss_max": float(np.max(reg_loss_list)),
        "MAE_mean": float(np.mean(reg_mae_list)),
        "RMSE_mean": float(np.mean(reg_rmse_list)),
        "MAPE_mean_%": float(np.mean(reg_mape_list)),
    })

for start in range(0, n - window + 1, stride):
    end = start + window

    xw = torch.tensor(Xn2[start:end], dtype=torch.float32)
    obs_w = {k: (v[start:end] if isinstance(v, np.ndarray) else v) for k, v in obs2.items()}

    Pf_w = obs_w["Pf"]
    TDSf_w = obs_w["TDSf"]
    Jw_w = obs_w["Jw_obs"]

    # ---- 门控1:jump gating ----
    if (Pf_w.max() - Pf_w.min()) > pf_jump_th or \
       (TDSf_w.max() - TDSf_w.min()) > tds_jump_th or \
       (Jw_w.max() - Jw_w.min()) > jw_jump_th:
        print(f"[sync-skip] reg={regime_id} window {start}:{end}  (jump gating)")
        consec_dist_fail = 0
        continue

    # ---- 门控2:distribution gating(z-score)----
    z_pf = abs(float(Pf_w.mean()) - mu_pf) / sd_pf_eff
    z_tds = abs(float(TDSf_w.mean()) - mu_tds) / sd_tds_eff
    z_jw = abs(float(Jw_w.mean()) - mu_jw) / sd_jw_eff

    if z_pf > z_th_pf or z_tds > z_th_tds or z_jw > z_th_jw:
        consec_dist_fail += 1
        print(f"[sync-skip] reg={regime_id} window {start}:{end}  "
              f"(dist gating: z_pf={z_pf:.2f}, z_tds={z_tds:.2f}, z_jw={z_jw:.2f})  "
              f"fail={consec_dist_fail}/{reset_patience}")

        if consec_dist_fail >= reset_patience:
            flush_regime_summary(regime_id)

            regime_id += 1
            mu_pf = float(Pf_w.mean()); sd_pf = float(Pf_w.std() + 1e-8)
            mu_tds = float(TDSf_w.mean()); sd_tds = float(TDSf_w.std() + 1e-8)
            mu_jw = float(Jw_w.mean()); sd_jw = float(Jw_w.std() + 1e-8)

            sd_pf_eff = max(sd_pf, sd_floor_pf)
            sd_tds_eff = max(sd_tds, sd_floor_tds)
            sd_jw_eff = max(sd_jw, sd_floor_jw)

            reg_mae_list, reg_rmse_list, reg_mape_list, reg_loss_list = [], [], [], []
            print(f"[baseline-a2:reset] reg={regime_id}  mu_pf={mu_pf:.4f} sd_pf={sd_pf:.4f}(eff={sd_pf_eff:.4f})  "
                  f"mu_tds={mu_tds:.2f} sd_tds={sd_tds:.2f}(eff={sd_tds_eff:.2f})  "
                  f"mu_jw={mu_jw:.2f} sd_jw={sd_jw:.2f}(eff={sd_jw_eff:.2f})")
            consec_dist_fail = 0

        continue

    # ---- 通过门控:同步更新 ----
    consec_dist_fail = 0
    for _ in range(steps):
        opt2.zero_grad()
        L, _ = loss_terms(model, xw, obs_w, lambdas=lambdas)
        L.backward()
        torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=clip_norm)
        opt2.step()

    mae, rmse, mape = metrics_jw(model, xw, obs_w)
    reg_mae_list.append(mae); reg_rmse_list.append(rmse); reg_mape_list.append(mape)
    reg_loss_list.append(float(L.item()))

    print(f"[sync] reg={regime_id} window {start}:{end}  L={L.item():.4e}  "
          f"MAE={mae:.3f} RMSE={rmse:.3f} MAPE={mape:.2f}%")

    # 记录 A/B 时序(窗口内每个点)
    Jw_pred, Js_pred, A_pred, B_pred = predict_ab(model, xw)
    t_sec = obs_w["t_sec"].reshape(-1)

    for i in range(len(t_sec)):
        ts_records.append({
            "index_global": int(start + i),
            "regime_id": int(regime_id),
            "t_sec": float(t_sec[i]),
            "Pf": float(obs_w["Pf"][i]),
            "TDSf": float(obs_w["TDSf"][i]),
            "Jw_obs": float(obs_w["Jw_obs"][i]),
            "Jw_pred": float(Jw_pred[i, 0]),
            "A_pred": float(A_pred[i, 0]),
            "B_pred": float(B_pred[i, 0]),
        })

flush_regime_summary(regime_id)

# ====== 写文件(输出到 OUT_DIR)======
df_ts = pd.DataFrame(ts_records)
df_reg = pd.DataFrame(regime_records)

# 去重:窗口重叠,同 index 可能多次记录 → 保留最后一次(最新同步结果)
if not df_ts.empty:
    df_ts = df_ts.sort_values(["index_global"]).drop_duplicates("index_global", keep="last")
    # ✅按时间排序再画图/保存,避免折线“乱连”
    df_ts = df_ts.sort_values("t_sec")

ab_path = OUT_DIR / out_ab_csv
reg_path = OUT_DIR / out_regime_csv
df_ts.to_csv(ab_path, index=False, encoding="utf-8-sig")
df_reg.to_csv(reg_path, index=False, encoding="utf-8-sig")

print(f"[saved] {ab_path}")
print(f"[saved] {reg_path}")

# ====== 画图(输出到 OUT_DIR)======
if not df_ts.empty:
    A_path = OUT_DIR / out_plot_A
    B_path = OUT_DIR / out_plot_B

    plt.figure()
    plt.plot(df_ts["t_sec"].values, df_ts["A_pred"].values)
    plt.xlabel("t (sec)")
    plt.ylabel("A_pred")
    plt.title("Predicted A(t)")
    plt.tight_layout()
    plt.savefig(A_path, dpi=200)

    plt.figure()
    plt.plot(df_ts["t_sec"].values, df_ts["B_pred"].values)
    plt.xlabel("t (sec)")
    plt.ylabel("B_pred")
    plt.title("Predicted B(t)")
    plt.tight_layout()
    plt.savefig(B_path, dpi=200)

    print(f"[saved] {A_path}")
    print(f"[saved] {B_path}")

=========================

9) 主流程

=========================

def main():
df1 = pd.read_excel(PATH_A1, sheet_name=SHEET_A1)
df2 = pd.read_excel(PATH_A2, sheet_name=SHEET_A2)

X1, obs1 = build_xy_a1(df1)
X2, obs2 = build_xy_a2_primary(df2)

X1n, mu1, sd1 = standardize(X1)
X2n = (X2 - mu1) / sd1

print("X1 dtype:", X1.dtype, "sample:", X1[0])
print("X2 dtype:", X2.dtype, "sample:", X2[0])
print(f"Output folder: {OUT_DIR}")

model = PINNNet(in_dim=5, hidden=64)

# 你可以调权重:更重视物理就把 lam_flux / lam_salt / lam_tds 调大
lambdas = (1.0, 1.0, 1.0, 1.0)

# 1) 离线训练
train_offline(model, X1n, obs1, epochs=2000, lr=1e-3, lambdas=lambdas)

# 2) 在线同步 + 导出
sync_online(
    model, X2n, obs2,
    window=200, stride=50,
    steps=200, lr=5e-5,
    lambdas=lambdas,
    clip_norm=1.0,
    pf_jump_th=0.30, tds_jump_th=2000.0, jw_jump_th=80.0,
    baseline_n=1000,
    z_th_pf=6.0, z_th_tds=6.0, z_th_jw=6.0,
    sd_floor_pf=0.02, sd_floor_tds=30.0, sd_floor_jw=5.0,
    reset_patience=6,
)

if name == "main":
main()

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

    Milestone

    No milestone

    Relationships

    None yet

    Development

    No branches or pull requests

    Issue actions