Files
dpb/backend/scripts/breeding_stats/mtblup.py
T

385 lines
16 KiB
Python
Raw Normal View History

"""动物模型 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,
}