File size: 4,793 Bytes
47c4bf8
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
322ebce
 
 
 
 
 
 
47c4bf8
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
"""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)]