Files
dpb/backend/scripts/breeding_stats/mtblup.py
T
34047007@qq.com b95053c52c init: 初始化 dpb 桃育种系统代码库
前后端 + 后端 FastAPI 全量源码、部署脚本与文档。
2026-08-06 00:17:49 +08:00

385 lines
16 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""动物模型 MT-BLUP(多性状 BLUP / 遗传相关估计)。
仅依赖 numpy,复用 blup 的系谱 A/A⁻¹、单性状剖面 REML 与黄金分割一维最大化。
两条路径:
- solve_bivariate:成对双性状(降维:Va/Ve 取单性状 REML,仅对 ρ 做 1-D 剖面最大化);
- solve_multi:全多变量 EM-REML(一次估计完整 G0⊗A,替代逐对,供 Smith-Hazel 指数)。
模型(成对双性状):
y = [y1; y2]Var(u) = G0 ⊗ AG0 = [[Va1, ρ·√(Va1·Va2)], [·, Va2]]
Var(e) = diag(Ve1·I, Ve2·I)。
降维策略(避免全多变量 REML 的收敛与识别风险):
Va1/Ve1、Va2/Ve2 直接取各性状单性状 blup.solve 的 REML 估计(= 既有 herit 来源),
仅对遗传相关 ρ ∈ [-0.99, 0.99] 做粗网格预扫 + 黄金分割最大化精确 REML 对数似然
V = Z(G0⊗A)Z' + Rslogdet + y'Py,与 _solve_gxe 模板一致)。
输出 σa12 = ρ·√(Va1·Va2)、r_g = ρ,供 Smith-Hazel 指数 G 矩阵非对角替换 Calo 近似。
单对不收敛 / 单性状求解失败 → converged=False + warning(调用方回退 Calo 并记审计)。
"""
from __future__ import annotations
import math
import numpy as np
from scripts.breeding_stats import blup
ENGINE_VERSION = "1.1.0" # 1.1.0: 全多变量 EM-REML solve_multiG0⊗A 整体估计)
_RHO_LO, _RHO_HI = -0.99, 0.99
_GRID_N = 21 # 粗网格预扫点数(防多峰取错局部峰)
_PAIR_KEYS = ["code1", "code2", "va1", "va2", "va12", "r_g", "ve1", "ve2",
"n_common", "n1", "n2", "n_individuals", "n_iter", "converged", "warning"]
def _pedigree_partition(pedigree: list[dict]):
"""系谱 -> (base, non_base, seen)。与 blup.solve 相同解析。"""
base: list[int] = []
non_base: list[tuple[int, int | None, int | None]] = []
seen: set[int] = set()
for rec in pedigree:
i = rec["individual"]
if i in seen:
raise ValueError(f"系谱个体重复: {i}")
seen.add(i)
d, s = rec.get("dam"), rec.get("sire")
if d is None and s is None:
base.append(i)
else:
non_base.append((i, d, s))
return base, non_base, seen
def _fail(code1: str, code2: str, msg: str) -> dict:
return {"code1": code1, "code2": code2, "va1": None, "va2": None, "va12": None,
"r_g": None, "ve1": None, "ve2": None, "n_common": 0, "n1": 0, "n2": 0,
"n_individuals": 0, "n_iter": 0, "converged": False, "warning": msg}
def solve_bivariate(pedigree: list[dict], phenos: dict[str, dict],
code1: str, code2: str, *, tol: float = blup.TOL) -> dict:
"""成对双性状 BLUP:估计 Va1/Va2/Va12(遗传相关 ρ)。
phenos: {code: {个体id: 均值}}(个体 id 空间与系谱一致,如 "t{tree_id}")。
"""
p1 = phenos.get(code1, {})
p2 = phenos.get(code2, {})
n1, n2 = len(p1), len(p2)
if n1 < 2 or n2 < 2:
return _fail(code1, code2, "任一性状表型个体不足 2 个,无法估计方差组分")
# 单性状方差分量(REML)——与 run_ablup 同一求解器,保证 herit 口径一致
try:
r1 = blup.solve(pedigree, p1, tol=tol)
r2 = blup.solve(pedigree, p2, tol=tol)
except Exception as e: # noqa: BLE001
return _fail(code1, code2, f"单性状求解失败: {e!s}")
va1, ve1 = float(r1["sigma_a"]), float(r1["sigma_e"])
va2, ve2 = float(r2["sigma_a"]), float(r2["sigma_e"])
if va1 <= 0 or va2 <= 0:
return _fail(code1, code2, "单性状加性方差为 0(数据不支持遗传方差),无法估计遗传相关")
base, non_base, seen = _pedigree_partition(pedigree)
for pid in list(p1) + list(p2):
if pid not in seen:
seen.add(pid)
base.append(pid)
if not seen:
return _fail(code1, code2, "系谱与表型均为空")
order, idx = blup._order_pedigree(base, non_base)
n = len(order)
obs1 = [ind for ind in order if ind in p1]
obs2 = [ind for ind in order if ind in p2]
m1, m2 = len(obs1), len(obs2)
ix1 = [idx[i] for i in obs1]
ix2 = [idx[i] for i in obs2]
y1 = np.array([float(p1[i]) for i in obs1], dtype=float)
y2 = np.array([float(p2[i]) for i in obs2], dtype=float)
y = np.concatenate([y1, y2])
X = np.zeros((m1 + m2, 2))
X[:m1, 0] = 1.0
X[m1:, 1] = 1.0
A = blup._build_ainv(order, len(base), {i: (d, s) for (i, d, s) in non_base})[0]
A11 = A[np.ix_(ix1, ix1)]
A12 = A[np.ix_(ix1, ix2)]
A22 = A[np.ix_(ix2, ix2)]
cross = float(np.sqrt(va1 * va2))
best = {"ll": -np.inf, "rho": 0.0}
n_eval = 0
def _ll(rho: float) -> float:
nonlocal n_eval
n_eval += 1
sa12 = rho * cross
V = np.zeros((m1 + m2, m1 + m2))
V[:m1, :m1] = va1 * A11 + ve1 * np.eye(m1)
V[:m1, m1:] = sa12 * A12
V[m1:, :m1] = sa12 * A12.T
V[m1:, m1:] = va2 * A22 + ve2 * np.eye(m2)
try:
VinvX = np.linalg.solve(V, X)
Vinvy = np.linalg.solve(V, y)
XtVinvX = X.T @ VinvX
XtVinvY = X.T @ Vinvy
yPy = float(y @ Vinvy) - float(XtVinvY @ np.linalg.solve(XtVinvX, XtVinvY))
_, lv = np.linalg.slogdet(V)
_, lx = np.linalg.slogdet(XtVinvX)
ll = -0.5 * (float(lv) + float(lx) + yPy)
except np.linalg.LinAlgError:
return -np.inf
if ll > best["ll"]:
best.update(ll=ll, rho=rho)
return ll
# 粗网格预扫 + 邻域黄金分割(防多峰取错局部峰)
grid = np.linspace(_RHO_LO, _RHO_HI, _GRID_N)
vals = [_ll(r) for r in grid]
best_i = int(np.argmax(vals))
lo = grid[max(best_i - 1, 0)]
hi = grid[min(best_i + 1, _GRID_N - 1)]
rho_opt, n_gold = blup._golden_max(_ll, lo, hi, tol=1e-4, max_iter=80)
if _ll(rho_opt) > best["ll"]:
best.update(ll=_ll(rho_opt), rho=rho_opt)
rho = best["rho"]
n_iter = _GRID_N + n_gold
warnings: list[str] = []
if abs(rho) >= _RHO_HI - 1e-6:
warnings.append(f"遗传相关 ρ 达边界 {rho:.3f}(接近完全{'正' if rho > 0 else '负'}相关,估计需谨慎)")
if n_iter >= blup.MAX_PROFILE_EVALS:
warnings.append("剖面 REML 未完全收敛(似然面极平/边界最优),结果采用已探明最优。")
common = set(obs1) & set(obs2)
return {
"code1": code1, "code2": code2,
"va1": float(va1), "va2": float(va2),
"va12": float(rho * cross),
"r_g": float(rho),
"ve1": float(ve1), "ve2": float(ve2),
"n_common": len(common), "n1": n1, "n2": n2,
"n_individuals": n, "n_iter": n_iter,
"converged": n_iter < blup.MAX_PROFILE_EVALS,
"warning": ("".join(warnings) if warnings else None),
}
def _fail_multi(codes: list[str], msg: str) -> dict:
return {"codes": codes, "G0": None, "G0_inv": None, "r_g": None, "ve": None,
"h2": None, "n_obs": None, "n_individuals": 0, "n_iter": 0,
"converged": False, "warning": msg, "engine_version": ENGINE_VERSION}
def solve_multi(pedigree: list[dict], phenos: dict[str, dict],
codes: list[str], *, tol: float = blup.TOL,
max_iter: int = 200, tol_em: float = 1e-4) -> dict:
"""全多变量 MT-BLUPEM-REML 估计完整 G0⊗A,替代逐对 bivariate)。
y = Xb + Zu + eu 按性状主序堆叠 [u_1;...;u_m](各 u_k 为全部 n 个体的育种值),
Var(u) = G0⊗AG0 m×m 遗传协方差,半正定),Var(e) = diag(ve_1·I, ..., ve_m·I)。
EM-REMLMeyer 1985):
- 初始 G0 = diag(各性状单性状 blup.solve REML va)、ve = 单性状 ve
- 迭代:MME(C·s=r) 求 β/u + C⁻¹ → 遗传 G0_ij = (û_i'A⁻¹û_j + tr(A⁻¹C_ij))/n、
残差 ve_i = (e_i'e_i + Σ_{k∈obs_i} C_ii[k,k]) / n_i
- G0 每次特征值截断投影半正定 + 对称化;|ΔREML LL| < tol_em 收敛;
- Aitken 逐元素外推加速(仅当外推点 REML LL 更高才接受,否则回退 EM 步)。
phenos: {code: {个体id: 均值}}(个体 id 空间与系谱一致,如 "t{tree_id}")。
输出 G0 全元素 + G0_invSmith-Hazel 指数) + r_g 相关矩阵 + ve + h2。
"""
if not codes or len(codes) < 2:
return _fail_multi(list(codes or []), "全 MT-BLUP 需至少 2 个性状")
codes = list(codes)
m = len(codes)
base, non_base, seen = _pedigree_partition(pedigree)
for code in codes:
for pid in phenos.get(code, {}):
if pid not in seen:
seen.add(pid)
base.append(pid)
if not seen:
return _fail_multi(codes, "系谱与表型均为空")
order, idx = blup._order_pedigree(base, non_base)
n = len(order)
parent_of = {i: (d, s) for (i, d, s) in non_base}
A, Ainv = blup._build_ainv(order, len(base), parent_of)
Xlist: list[np.ndarray] = []
Zlist: list[np.ndarray] = []
ylist: list[np.ndarray] = []
obslist: list[list[int]] = []
va0: list[float] = []
ve0: list[float] = []
for code in codes:
pk = phenos.get(code, {})
obs_k = [ind for ind in order if ind in pk]
if len(obs_k) < 2:
return _fail_multi(codes, f"性状 {code} 表型个体不足 2,无法估计方差组分")
obslist.append(obs_k)
ylist.append(np.array([float(pk[i]) for i in obs_k], dtype=float))
Zk = np.zeros((len(obs_k), n))
for kk, ind in enumerate(obs_k):
Zk[kk, idx[ind]] = 1.0
Zlist.append(Zk)
Xlist.append(np.ones((len(obs_k), 1)))
rk = blup.solve(pedigree, pk, tol=tol)
va0.append(float(rk["sigma_a"]))
ve0.append(float(rk["sigma_e"]))
if any(v <= 0 for v in va0):
return _fail_multi(codes, "有性状加性方差为 0,全多变量 REML 无法启动(请回退成对 bivariate")
n_obs_list = [len(o) for o in obslist]
N = sum(n_obs_list)
Xbig = np.zeros((N, m))
Zbig = np.zeros((N, m * n))
yvec = np.zeros(N)
off = 0
for k in range(m):
mk = n_obs_list[k]
Xbig[off:off + mk, k] = 1.0
Zbig[off:off + mk, k * n:(k + 1) * n] = Zlist[k]
yvec[off:off + mk] = ylist[k]
off += mk
def _reml_ll(G0c: np.ndarray, vec: np.ndarray) -> float:
ZAZ = np.zeros((N, N))
for i in range(m):
for j in range(m):
offi = sum(n_obs_list[:i])
offj = sum(n_obs_list[:j])
ZAZ[offi:offi + n_obs_list[i], offj:offj + n_obs_list[j]] = (
G0c[i, j] * (Zlist[i] @ A @ Zlist[j].T))
V = ZAZ + np.diag(np.repeat(vec, n_obs_list))
try:
VinvX = np.linalg.solve(V, Xbig)
Vinvy = np.linalg.solve(V, yvec)
XtVinvX = Xbig.T @ VinvX
XtVinvY = Xbig.T @ Vinvy
yPy = float(yvec @ Vinvy) - float(XtVinvY @ np.linalg.solve(XtVinvX, XtVinvY))
_, lv = np.linalg.slogdet(V)
_, lx = np.linalg.slogdet(XtVinvX)
return -0.5 * (float(lv) + float(lx) + yPy)
except np.linalg.LinAlgError:
return -np.inf
def _project_psd(mat: np.ndarray) -> tuple[np.ndarray, bool]:
mat = 0.5 * (mat + mat.T)
w, V = np.linalg.eigh(mat)
floor = 1e-6 * max(float(w[-1]), 1e-6)
clipped = bool(w[0] <= floor)
return (V * np.maximum(w, floor)) @ V.T, clipped
G0 = np.diag(va0)
ve = np.array(ve0, dtype=float)
warnings: list[str] = []
converged = False
ll = -np.inf
prev_ll: float | None = None
ll_gain = 1.0
it = 0
theta_prev2: np.ndarray | None = None
theta_prev1: np.ndarray | None = None
def _em_step(par_g0: np.ndarray, par_ve: np.ndarray) -> tuple[np.ndarray, np.ndarray, float]:
Rinv_diag = np.repeat(1.0 / par_ve, n_obs_list)
G0inv = np.linalg.inv(par_g0)
C = np.zeros((m + m * n, m + m * n))
C[:m, :m] = Xbig.T @ (Rinv_diag[:, None] * Xbig)
XtRinvZ = Xbig.T @ (Rinv_diag[:, None] * Zbig)
C[:m, m:] = XtRinvZ
C[m:, :m] = XtRinvZ.T
C[m:, m:] = Zbig.T @ (Rinv_diag[:, None] * Zbig) + np.kron(G0inv, Ainv)
rhs = np.concatenate([Xbig.T @ (Rinv_diag * yvec), Zbig.T @ (Rinv_diag * yvec)])
try:
sol = np.linalg.solve(C, rhs)
Cinv = np.linalg.inv(C)
except np.linalg.LinAlgError:
sol = np.linalg.pinv(C) @ rhs
Cinv = np.linalg.pinv(C)
u = sol[m:]
Cuu = Cinv[m:, m:]
G0n = np.zeros((m, m))
for i in range(m):
ui = u[i * n:(i + 1) * n]
for j in range(m):
uj = u[j * n:(j + 1) * n]
Cij = Cuu[i * n:(i + 1) * n, j * n:(j + 1) * n]
G0n[i, j] = (float(ui @ Ainv @ uj) + float(np.sum(Ainv * Cij))) / n
G0n, proj = _project_psd(G0n)
if proj:
warnings.append("遗传协方差 G0 非半正定,已特征值截断投影到最近半正定阵")
ve_n = np.zeros(m)
for k in range(m):
e_k = ylist[k] - float(sol[k]) - Zlist[k] @ u[k * n:(k + 1) * n]
obs_pos = np.array([idx[i] for i in obslist[k]])
# tr(PEV_e_k) = tr(W_k C⁻¹_kk W_k')W_k=[X_k,Z_k]X 为截距列、Z 为逐个体指示)
# = n_k·C⁻¹[βk,βk] + 2·Σ_obs C⁻¹[βk,u_k] + Σ_obs C⁻¹[u_k,u_k]
block_ku = Cinv[k, m + k * n: m + (k + 1) * n]
block_uu_k = Cinv[m + k * n: m + (k + 1) * n, m + k * n: m + (k + 1) * n]
trace_r = (n_obs_list[k] * float(Cinv[k, k])
+ 2.0 * float(np.sum(block_ku[obs_pos]))
+ float(np.sum(block_uu_k[obs_pos, obs_pos])))
ve_n[k] = (float(e_k @ e_k) + trace_r) / n_obs_list[k]
ve_n = np.clip(ve_n, blup._FLOOR, None)
return G0n, ve_n, _reml_ll(G0n, ve_n)
for it in range(1, max_iter + 1):
G0new, ve_new, new_ll = _em_step(G0, ve)
if new_ll == -np.inf:
continue
theta_new = np.concatenate([G0new.ravel(), ve_new])
# Aitken 加速:逐元素外推 θ∞≈θₜ+d₁/(1−ρ),仅当外推点 REML LL 更高才接受,
# 缓解 EM 在高相关多性状下收敛缓慢(trait 方差先塌缩后缓慢恢复的困境)。
if (it >= 4 and theta_prev2 is not None and theta_prev1 is not None):
d1 = theta_new - theta_prev1
d2 = theta_prev1 - theta_prev2
denom = np.where(np.abs(d2) > 1e-12, d2, 1.0)
rho = np.where(np.abs(d2) > 1e-12, d1 / denom, 0.0)
rho = np.clip(rho, -0.999, 0.999)
theta_acc = theta_new + d1 / (1.0 - rho)
G0a, _ = _project_psd(0.5 * (theta_acc[:m * m].reshape(m, m)
+ theta_acc[:m * m].reshape(m, m).T))
vea = np.clip(theta_acc[m * m:], blup._FLOOR, None)
ll_acc = _reml_ll(G0a, vea)
if ll_acc > new_ll and ll_acc != -np.inf:
G0new, ve_new, new_ll = G0a, vea, ll_acc
theta_prev2, theta_prev1 = theta_prev1, theta_new
ll_gain = abs(new_ll - prev_ll) if prev_ll is not None else 1.0
G0, ve = G0new, ve_new
ll = new_ll
prev_ll = new_ll
if ll_gain < tol_em:
converged = True
break
if not converged:
warnings.append(f"EM-REML {max_iter} 次迭代未完全收敛(|ΔLL|={ll_gain:.2e}),结果采用已探明最优")
r_g = np.zeros((m, m))
for i in range(m):
for j in range(m):
if G0[i, i] > 0 and G0[j, j] > 0:
r_g[i, j] = G0[i, j] / math.sqrt(G0[i, i] * G0[j, j])
h2 = [G0[k, k] / (G0[k, k] + ve[k]) if (G0[k, k] + ve[k]) > 0 else None
for k in range(m)]
G0i = np.linalg.inv(G0)
return {
"codes": codes,
"G0": [[round(float(G0[i, j]), 6) for j in range(m)] for i in range(m)],
"G0_inv": [[round(float(G0i[i, j]), 6) for j in range(m)] for i in range(m)],
"r_g": [[round(float(r_g[i, j]), 4) for j in range(m)] for i in range(m)],
"ve": [round(float(v), 6) for v in ve],
"h2": [round(float(h), 4) if h is not None else None for h in h2],
"n_obs": n_obs_list,
"n_individuals": n,
"n_iter": it,
"converged": converged,
"warning": "".join(warnings) if warnings else None,
"engine_version": ENGINE_VERSION,
}