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()
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")
=========================
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)
def build_xy_a2_primary(df: pd.DataFrame):
df = df.copy()
=========================
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)
=========================
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
@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)
=========================
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,
):
opt2 = torch.optim.Adam(model.parameters(), lr=lr)
n = Xn2.shape[0]
=========================
9) 主流程
=========================
def main():
df1 = pd.read_excel(PATH_A1, sheet_name=SHEET_A1)
df2 = pd.read_excel(PATH_A2, sheet_name=SHEET_A2)
if name == "main":
main()