Rootscope / rootscope /classify_v4.py
ct-tranchau's picture
RootScope v4: manuscript model (LightGBM per round, 3 seeds, layer + soft-neighbor context, radial prior)
322ebce verified
Raw History Blame Contribute Delete
4.79 kB
"""RootScope v4 -- tang phan loai.
Khac v2 o ba diem, va ca ba deu nam o day chu khong o phan trich dac trung:
1. Hang xom dang MEM: thay vi chot nhan cua hang xom roi dua vao model, moi huong
dong gop ca vector xac suat 9 lop (36 cot) + 6 co dong thuan.
2. Dac trung LOP: khoang cach tu moi te bao toi cac lop mo (epidermis, exodermis,
endodermis, pericycle) ma chinh model da du doan o vong truoc, do theo ban kinh
cuc bo va tinh bang duong kinh te bao.
3. Radial-order prior: uu tien thu tu xuyen tam dung giai phau, ap sau vong cuoi.
Cac ham tinh dac trung nam o context.py va duoc CAT NGUYEN VAN tu code nghien cuu.
"""
from pathlib import Path
import joblib
import numpy as np
import pandas as pd
from . import context as ctx
from . import _eval_honest as _eh
from .extract_features import build_cell_adjacency
from .radial_order import correct
CLASSES = ["cortex", "endodermis", "epidermis", "exodermis", "pericycle",
"phloem", "root_cap", "stele", "xylem"]
PRIOR = (0.3, 2.0) # (s, lambda) cua model san xuat
def load_models(model_dir, seeds=None):
"""Nap cac goi LightGBM v4. Moi goi chua round_models = model cua tung vong."""
model_dir = Path(model_dir)
out = []
for p in sorted(model_dir.glob("lgbm_s*.joblib")):
s = int(p.stem.split("_s")[-1])
if seeds is None or s in seeds:
b = joblib.load(p)
if "round_models" not in b:
raise ValueError(
f"{p.name} khong co 'round_models'. Goi nay chi luu model vong cuoi "
f"nen khong chay lai duoc vong lap. Dung goi san xuat.")
out.append((s, b))
if not out:
raise FileNotFoundError(f"khong co lgbm_s*.joblib trong {model_dir}")
f0 = list(out[0][1]["fcols"])
for s, b in out:
if list(b["fcols"]) != f0:
raise ValueError(f"seed {s} co bo feature khac")
if list(b["classes"]) != CLASSES:
raise ValueError(f"seed {s} co thu tu lop khac")
return out
def classify(df_base, masks, source_file="image.tif", model_dir=None, seeds=None,
max_rounds=6, prior=PRIOR, verbose=True):
"""Tra (nhan, xac suat, nhan vong 1) cho tung dong cua df_base."""
bundles = load_models(model_dir, seeds)
fcols = list(bundles[0][1]["fcols"])
geom = df_base.copy()
geom["source_file"] = source_file
geom = geom.reset_index(drop=True)
# Nap thang adjacency vao cache cua _eval_honest: khoi phai ghi mask ra dia
# roi doc lai nhu duong danh gia trong du an.
_eh._ADJ[source_file] = build_cell_adjacency(masks)
ctx.MASK = "<in-memory>"
ctx.CL = CLASSES
# The 4 nearest-neighbour cell-type columns are model inputs that exist only
# as placeholders before round 1 (-1 = unknown); stage_features does not
# produce them, so create them here. fill_neighbors overwrites them from round 2.
for c in _eh.NEIGHBOR_CT_COLS:
if c not in geom.columns:
geom[c] = -1.0
missing = [c for c in fcols if c not in geom.columns
and c not in ctx.RINGCOLS and not c.startswith(("nbp_", "agree_"))]
if missing:
raise ValueError(f"thieu {len(missing)} cot dac trung, vd {missing[:6]}")
mats = ctx.dir_mats(geom)
P1s, PFs = [], []
for s, b in bundles:
prev = prevP = None
for rnd in range(1, max_rounds + 1):
cur = ctx._blank(geom) if rnd == 1 else ctx.fill_neighbors(geom, prev, ctx.MASK)
cur = pd.concat([cur.reset_index(drop=True),
ctx.soft_feats(prevP, mats, CLASSES),
ctx.ring_feats(geom, prev, "q", prevP)], axis=1)
m_r, s_r = b["round_models"][min(rnd, len(b["round_models"])) - 1]
P = m_r.predict_proba(s_r.transform(cur[fcols].values.astype(np.float32)))
if rnd == 1:
P1s.append(P.astype(np.float64))
names = np.array(CLASSES)[P.argmax(1)]
chg = -1 if prev is None else int((names != prev).sum())
if verbose:
print(f" seed {s} round {rnd}: {chg} labels changed", flush=True)
prev, prevP = names, P
if chg == 0:
break
PFs.append(P.astype(np.float64))
P1 = np.mean(P1s, 0)
P = np.mean(PFs, 0)
if prior and prior[1]:
cy = geom.centroid_y.values.astype(float)
cx = geom.centroid_x.values.astype(float)
cid = geom.cell_id.values
look = {int(i): (y, x) for i, y, x in zip(cid, cy, cx)}
P = correct(P, np.full(len(geom), source_file), cid, CLASSES,
geom={source_file: (cy, cx, look)}, s=prior[0], lam=prior[1])
return np.array(CLASSES)[P.argmax(1)], P, np.array(CLASSES)[P1.argmax(1)]