"""动物模型 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 ⊗ A,G0 = [[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' + R,slogdet + 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_multi(G0⊗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-BLUP(EM-REML 估计完整 G0⊗A,替代逐对 bivariate)。 y = Xb + Zu + e:u 按性状主序堆叠 [u_1;...;u_m](各 u_k 为全部 n 个体的育种值), Var(u) = G0⊗A(G0 m×m 遗传协方差,半正定),Var(e) = diag(ve_1·I, ..., ve_m·I)。 EM-REML(Meyer 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_inv(Smith-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, }