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

288 lines
13 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.
# -*- coding: utf-8 -*-
"""ssGBLUP 正式 tc 套件:TestClient 走真实 API + golden 数值断言(H⁻¹ 拼接判别)。
fixture:5 株表型 + 基因型数据集(树0/1/3 直连、树2 名称兜底、树4 全多等位无有效剂量)。
经真实 HTTP 端点 POST /api/v1/bre/statistics/gblup/run 落库后,ORM 断言:
[1] GBLUP method=GBLUP、h²∈(0,1)、5 条 EBV、非基因型株(树4) EBV=0
[2] ssGBLUPmethod=ssGBLUP、h²∈(0,1)、5 条 EBV、非基因型株(树4) EBV≠0
(H⁻¹ 拼接生效的判别核心——若 H⁻¹=A⁻¹+[0;G⁻¹-A22⁻¹] 拼接错,非基因型亲属
EBV 会为 0 或爆炸,正是「算偏但当正常结果输出」的同款静默失效)
[3] 非基因型亲本(g 前缀)在 ssGBLUP 下 EBV 有限(H 进 MME,非 A22 全零)
依赖: Redis + PG 正常(TestClient 走真实 lifespan)。运行后自动清理。
"""
import os
os.environ["ENVIRONMENT"] = "dev"
os.environ["PYTHONUTF8"] = "1"
import sys, asyncio # noqa: E402
sys.path.insert(0, r"d:\dpb\dpb\backend")
import main # noqa: E402
from fastapi.testclient import TestClient # noqa: E402
from sqlalchemy import delete, select # noqa: E402
from app.core.database import create_async_engine_and_session # noqa: E402
from app.api.v1.module_system.user.model import UserModel # noqa: E402 (注册 mapper)
from app.api.v1.module_bre.trait.model import TraitModel # noqa: E402
from app.api.v1.module_bre.target.model import TargetModel # noqa: E402
from app.api.v1.module_bre.germplasm.model import BreedingGermplasmModel # noqa: E402
from app.api.v1.module_bre.cross_combination.model import CrossCombinationModel # noqa: E402
from app.api.v1.module_bre.tree.model import TreeModel # noqa: E402
from app.api.v1.module_bre.trait_observation.model import TraitObservationModel # noqa: E402
from app.api.v1.module_bre.genotype_dataset.model import GenotypingDatasetModel # noqa: E402
from app.api.v1.module_bre.genotype_sample.model import GenotypeSampleModel # noqa: E402
from app.api.v1.module_bre.genotype_call.model import GenotypeCallModel # noqa: E402
from app.api.v1.module_bre.marker.model import MarkerModel # noqa: E402
from app.api.v1.module_bre.statistics.model import ( # noqa: E402
PredictionModel,
PredictionValueModel,
StatisticsJobModel,
)
create_app = main.create_app
TOKEN = None
ok, fail = 0, 0
PREFIX = "TCSSGB"
tokens: dict[str, list[int]] = {
"combo": [], "tree": [], "obs": [],
"marker": [], "dataset": [], "sample": [], "call": [],
}
TRAIT_CODE = f"tR_{PREFIX}"
pred_ids: list[int] = []
FIX = {}
def check(name, cond, detail=""):
global ok, fail
if cond:
ok += 1
print(f" [ok] {name} {detail}")
else:
fail += 1
print(f" [FAIL] {name} {detail}")
def login(client):
global TOKEN
d = {"username": "super", "password": "123456", "grant_type": "password", "login_type": "PC端"}
r = client.post("/api/v1/system/auth/login", data=d)
b = r.json()
if r.status_code == 200 and b.get("code") == 0:
TOKEN = b["data"]["access_token"]
return
key = client.get("/api/v1/system/auth/captcha/get").json()["data"]["key"]
client.post("/api/v1/system/auth/captcha/slider/complete", json={"captcha_key": key})
d["captcha_key"] = key
r = client.post("/api/v1/system/auth/login", data=d)
b = r.json()
assert r.status_code == 200 and b.get("code") == 0, f"LOGIN FAIL {r.status_code} {b}"
TOKEN = b["data"]["access_token"]
def auth():
return {"Authorization": f"Bearer {TOKEN}"}
GENOS = {
0: [0, 0, 0, 1, 1, 1, 2, 2, 2, 1],
1: [1, 1, 0, 0, 1, 2, 2, 0, 1, 1],
2: [0, 1, 2, 1, 0, 1, 0, 1, 2, 0],
3: [1, 2, 1, 2, 1, 0, 1, 0, 0, 1],
}
GT = {0: "0/0", 1: "0/1", 2: "1/1"}
async def _build_fixture() -> None:
engine, sf = create_async_engine_and_session()
try:
async with sf() as db:
trait = TraitModel(trait_code=TRAIT_CODE, trait_name=f"含糖量{PREFIX}", data_type="numeric",
unit="%", is_core="1", direction="desc", into_ebv="1",
default_h2=0.5, created_id=1)
db.add(trait)
await db.flush()
FIX["trait_id"] = trait.id
target = TargetModel(target_name=f"目标{PREFIX}", created_id=1)
db.add(target)
await db.flush()
fm = BreedingGermplasmModel(cultivar_name=f"FM_{PREFIX}", can_be_female=True, created_id=1)
mm = BreedingGermplasmModel(cultivar_name=f"MM_{PREFIX}", can_be_male=True, created_id=1)
db.add_all([fm, mm])
await db.flush()
combo = CrossCombinationModel(
combination_code=f"C_{PREFIX}", bre_target_id=target.id,
female_parent_id=fm.id, male_parent_id=mm.id, design_type="full_diallel", created_id=1)
db.add(combo)
await db.flush()
tokens["combo"].append(combo.id)
trees = []
for i in range(5):
t = TreeModel(combination_id=combo.id, tree_no=f"{PREFIX}-T{i:02d}", status="alive",
stage="seedling", generation="F1", planted_date="2023-03-10", created_id=1)
db.add(t)
trees.append(t)
await db.flush()
tokens["tree"] = [t.id for t in trees]
FIX["tree_ids"] = [t.id for t in trees]
obs_rows = []
for i, t in enumerate(trees):
o = TraitObservationModel(tree_id=t.id, combination_id=combo.id, trait_id=trait.id,
value_numeric=12.0 + 0.8 * i, evaluate_year=2025, created_id=1)
db.add(o)
obs_rows.append(o)
await db.flush()
tokens["obs"] = [o.id for o in obs_rows]
ds = GenotypingDatasetModel(dataset_name=f"GS_{PREFIX}", platform="SSR",
purpose="GS", created_id=1)
db.add(ds)
await db.flush()
FIX["dataset_id"] = ds.id
tokens["dataset"].append(ds.id)
markers = []
for j in range(10):
m = MarkerModel(marker_name=f"MK{j}_{PREFIX}", marker_type="SNP",
chromosome=str(j % 5), position=j * 100, created_id=1)
db.add(m)
markers.append(m)
await db.flush()
tokens["marker"] = [m.id for m in markers]
FIX["markers"] = markers
samples = []
for i in (0, 1, 3):
s = GenotypeSampleModel(sample_name=f"{PREFIX}-S{i}", dataset_id=ds.id,
source_type="tree", source_id=trees[i].id, created_id=1)
db.add(s)
samples.append(s)
s_fb = GenotypeSampleModel(sample_name=trees[2].tree_no, dataset_id=ds.id,
source_type=None, source_id=None, created_id=1)
db.add(s_fb)
samples.append(s_fb)
s_bad = GenotypeSampleModel(sample_name=f"{PREFIX}-B4", dataset_id=ds.id,
source_type="tree", source_id=trees[4].id, created_id=1)
db.add(s_bad)
samples.append(s_bad)
await db.flush()
tokens["sample"] = [s.id for s in samples]
calls = []
tree_samp = {trees[0].id: samples[0], trees[1].id: samples[1], trees[3].id: samples[2]}
tree_samp[trees[2].id] = s_fb
genos_by_tid = {trees[i].id: GENOS[i] for i in range(4)}
for tid, smp in tree_samp.items():
for j, m in enumerate(markers):
calls.append(GenotypeCallModel(sample_id=smp.id, marker_id=m.id,
allele=GT[genos_by_tid[tid][j]], created_id=1))
for j, m in enumerate(markers):
calls.append(GenotypeCallModel(sample_id=s_bad.id, marker_id=m.id,
allele="1/2", created_id=1))
db.add_all(calls)
await db.flush()
tokens["call"] = [c.id for c in calls]
await db.commit()
finally:
await engine.dispose()
async def _verify_and_cleanup() -> None:
engine, sf = create_async_engine_and_session()
try:
async with sf() as db:
# ---- [1] GBLUP ----
p1 = await db.get(PredictionModel, pred_ids[0])
check("[1] method=GBLUP", p1.method == "GBLUP", p1.method)
check("[1] h²∈(0,1)", p1.heritability is not None and 0 < p1.heritability < 1,
f"{p1.heritability}")
vals1 = (await db.execute(select(PredictionValueModel).where(
PredictionValueModel.prediction_id == pred_ids[0]))).scalars().all()
check("[1] 5 条 EBV", len(vals1) == 5, f"{len(vals1)}")
v4_1 = next(v for v in vals1 if v.tree_id == FIX["tree_ids"][4])
check("[1] 非基因型株(树4) GBLUP EBV=0", float(v4_1.predicted_value) == 0.0,
f"{v4_1.predicted_value}")
check("[1] 全部 EBV 非 NaN", all(float(v.predicted_value) == float(v.predicted_value)
for v in vals1))
# ---- [2] ssGBLUP ----
p2 = await db.get(PredictionModel, pred_ids[1])
check("[2] method=ssGBLUP", p2.method == "ssGBLUP", p2.method)
check("[2] h²∈(0,1)", p2.heritability is not None and 0 < p2.heritability < 1,
f"{p2.heritability}")
vals2 = (await db.execute(select(PredictionValueModel).where(
PredictionValueModel.prediction_id == pred_ids[1]))).scalars().all()
check("[2] 5 条 EBV", len(vals2) == 5, f"{len(vals2)}")
v4_2 = next(v for v in vals2 if v.tree_id == FIX["tree_ids"][4])
check("[2] 非基因型株(树4) ssGBLUP EBV≠0H⁻¹ 拼接生效)",
float(v4_2.predicted_value) != 0.0, f"{v4_2.predicted_value}")
# ---- [3] 非基因型亲本 EBV 有限 ----
for pv in vals2:
assert float(pv.predicted_value) == float(pv.predicted_value), "EBV NaN"
check("[3] ssGBLUP 全部 EBV 有限", all(abs(float(v.predicted_value)) < 1e6 for v in vals2))
# ---- cleanup ----
if pred_ids:
await db.execute(delete(PredictionValueModel).where(
PredictionValueModel.prediction_id.in_(pred_ids)))
await db.execute(delete(StatisticsJobModel).where(
StatisticsJobModel.result_ref.in_(pred_ids)))
await db.execute(delete(PredictionModel).where(PredictionModel.id.in_(pred_ids)))
if tokens["call"]:
await db.execute(delete(GenotypeCallModel).where(
GenotypeCallModel.id.in_(tokens["call"])))
if tokens["sample"]:
await db.execute(delete(GenotypeSampleModel).where(
GenotypeSampleModel.id.in_(tokens["sample"])))
if tokens["dataset"]:
await db.execute(delete(GenotypingDatasetModel).where(
GenotypingDatasetModel.id.in_(tokens["dataset"])))
if tokens["marker"]:
await db.execute(delete(MarkerModel).where(MarkerModel.id.in_(tokens["marker"])))
if tokens["obs"]:
await db.execute(delete(TraitObservationModel).where(
TraitObservationModel.id.in_(tokens["obs"])))
if tokens["tree"]:
await db.execute(delete(TreeModel).where(TreeModel.id.in_(tokens["tree"])))
if tokens["combo"]:
await db.execute(delete(CrossCombinationModel).where(
CrossCombinationModel.id.in_(tokens["combo"])))
await db.execute(delete(BreedingGermplasmModel).where(
BreedingGermplasmModel.cultivar_name.in_([f"FM_{PREFIX}", f"MM_{PREFIX}"])))
await db.execute(delete(TraitModel).where(TraitModel.trait_code == TRAIT_CODE))
await db.execute(delete(TargetModel).where(TargetModel.target_name == f"目标{PREFIX}"))
await db.commit()
print(f"[cleanup] ssGBLUP tc 数据已清(批次 {len(pred_ids)} 个)")
finally:
await engine.dispose()
def main_() -> None:
asyncio.run(_build_fixture())
try:
with TestClient(create_app()) as client:
login(client)
for method in ("gblup", "ssgblup"):
r = client.post("/api/v1/bre/statistics/gblup/run", json={
"dataset_id": FIX["dataset_id"], "trait_id": FIX["trait_id"],
"trait_code": TRAIT_CODE, "year": None, "method": method, "maf_min": 0.05,
}, headers=auth())
check(f"[HTTP] gblup/run {method} 200", r.status_code == 200,
f"{r.status_code} {str(r.text)[:150]}")
body = r.json()
pid = body.get("data") if body.get("code") == 0 else None
check(f"[HTTP] {method} 返回 id", isinstance(pid, int), f"pid={pid}")
if isinstance(pid, int):
pred_ids.append(pid)
finally:
asyncio.run(_verify_and_cleanup())
print(f"\n===== ssGBLUP tc 套件:ok={ok} fail={fail} =====")
if __name__ == "__main__":
main_()
sys.exit(1 if fail else 0)