507 lines
23 KiB
Python
507 lines
23 KiB
Python
|
|
#!/home/xfr/.conda/envs/cbt/bin/python
|
|||
|
|
"""把標準化 grid 上的 segmentation mask 反向映射回 data_root 原始 CT grid。
|
|||
|
|
|
|||
|
|
座標參數來自 xfr_preprocess 記錄的 <vol_dir>/transform.json(每 level:
|
|||
|
|
裁切 box / standardize 翻軸 / ap_flip / 原 CT 幾何 / rotated 的 R、center、
|
|||
|
|
start)。純 index 空間映射(不依賴 standardized 輸出的物理 header),NN
|
|||
|
|
(order=0),輸出幾何 = 原始 CT,pixel type = 輸入 mask。
|
|||
|
|
|
|||
|
|
支援的 source grid(--source,預設由檔名自動判定):
|
|||
|
|
rotated <vol>/rotated/<L>_*.nii.gz 上的任何 mask(含 label map、
|
|||
|
|
自己跑的分割結果;整格同幾何)
|
|||
|
|
smd_resampled <vol>[/crop]/<L>_smd_resampled.nii.gz(0.5mm 裁切)
|
|||
|
|
binary_sdf <vol>[/crop]/<L>_binary_sdf.nii.gz
|
|||
|
|
roi <vol>[/crop]/<L>_roi.nii.gz
|
|||
|
|
binary_nn <vol>[/crop]/<L>_binary_nn.nii.gz
|
|||
|
|
binary <vol>[/crop]/<L>_binary.nii.gz(原解析度)
|
|||
|
|
smd <vol>[/crop]/<L>_smd.nii.gz(原解析度 SMD,裁切 +4 margin)
|
|||
|
|
([/crop]:新世代未旋轉檔在 crop/ 子資料夾、舊世代在頂層,自動判定)
|
|||
|
|
mask 尺寸須與該 source 檔一致(同 grid);level 取檔名前綴(L1..L5/T#..)。
|
|||
|
|
|
|||
|
|
用法:
|
|||
|
|
xfr_inverse_transform.py VOL_DIR MASK [--source auto] [--level auto]
|
|||
|
|
[--out PATH] [--orig CT] [--rebuild] [--chunk 16]
|
|||
|
|
|
|||
|
|
VOL_DIR 標準化 volume 目錄(含 transform.json 或可 --rebuild)
|
|||
|
|
MASK mask nifti(label / mask / 浮點分數皆可,dtype 保留)
|
|||
|
|
--rebuild transform.json 缺失時從 disk + data_root + metadata db 重建
|
|||
|
|
(box 由 disk 檔 origin 回復(精確;缺檔時 fallback native
|
|||
|
|
label 重算)、rotated R/center 重算並做正向 identity 驗證、
|
|||
|
|
fstart 由 disk 幾何回復、ap_flip 以骨頭/native label 對位 IoU
|
|||
|
|
檢測(db 值僅先驗——舊世代與新世代 disk 資料的 ap_flip 可能
|
|||
|
|
不同);結果寫回 transform.json)。舊世代(xfr-2 等)用這個。
|
|||
|
|
|
|||
|
|
輸出預設:<MASK>_to_original.nii.gz(與 MASK 同目錄)。
|
|||
|
|
"""
|
|||
|
|
|
|||
|
|
import argparse
|
|||
|
|
import logging
|
|||
|
|
import os
|
|||
|
|
import re
|
|||
|
|
import sys
|
|||
|
|
import time
|
|||
|
|
|
|||
|
|
import numpy as np
|
|||
|
|
import SimpleITK as sitk
|
|||
|
|
|
|||
|
|
from config.constant import LABEL_MAP
|
|||
|
|
from imaging.orientation import best_symmetry_plane, best_upper_endplate_plane
|
|||
|
|
from imaging.resample import resample_img
|
|||
|
|
from imaging.segmentation import _largest_cc_bbox
|
|||
|
|
from imaging.transforms import (SOURCES, build_volume_meta, img_geom,
|
|||
|
|
level_file_path, load_transform,
|
|||
|
|
mask_to_original, margined_box,
|
|||
|
|
resampled05_geom_from_original, save_transform,
|
|||
|
|
std_flip_axes_for_direction, transform_path)
|
|||
|
|
from visualization.res_bone_figure import compute_normalizing_rotation
|
|||
|
|
|
|||
|
|
from xfr_orig_labels import (_fstart_from_geometries, _forward_identity_check,
|
|||
|
|
_load_unrotated_template, _bone_mask,
|
|||
|
|
find_original_ct, load_ap_flip)
|
|||
|
|
|
|||
|
|
DATA_ROOT = '/mnt/1220/Public/dataset/Spine/CTSpine1K/data/'
|
|||
|
|
LABEL_ROOT = '/mnt/1220/Public/dataset/Spine/CTSpine1K/label/'
|
|||
|
|
LABEL_SUBDIRS = ('conlon', 'COVID-19', 'HNSCC-3DCT-RT_neck', 'Liver')
|
|||
|
|
|
|||
|
|
LEVEL_NAMES = set(LABEL_MAP.values())
|
|||
|
|
_LABEL_ID = {v: int(k) for k, v in LABEL_MAP.items()}
|
|||
|
|
|
|||
|
|
LOG_DIR = os.path.join(os.path.dirname(os.path.abspath(__file__)), 'logs')
|
|||
|
|
|
|||
|
|
logger = logging.getLogger('xfr_inverse_transform')
|
|||
|
|
|
|||
|
|
|
|||
|
|
def setup_tee(log_path):
|
|||
|
|
"""console 與 log 檔同時輸出(append、line-buffered)。"""
|
|||
|
|
class _Tee:
|
|||
|
|
def __init__(self, console, fh):
|
|||
|
|
self.console, self.fh, self.buf = console, fh, ''
|
|||
|
|
|
|||
|
|
def write(self, data):
|
|||
|
|
if not data:
|
|||
|
|
return
|
|||
|
|
self.buf += data
|
|||
|
|
while True:
|
|||
|
|
idx_n, idx_r = self.buf.find('\n'), self.buf.find('\r')
|
|||
|
|
idx = min([i for i in (idx_n, idx_r) if i != -1], default=-1)
|
|||
|
|
if idx == -1:
|
|||
|
|
break
|
|||
|
|
line, self.buf = self.buf[:idx], self.buf[idx + 1:]
|
|||
|
|
self.console.write(line + '\n')
|
|||
|
|
self.fh.write(line + '\n')
|
|||
|
|
|
|||
|
|
def flush(self):
|
|||
|
|
if self.buf:
|
|||
|
|
line, self.buf = self.buf, ''
|
|||
|
|
self.console.write(line + '\n')
|
|||
|
|
self.fh.write(line + '\n')
|
|||
|
|
self.console.flush()
|
|||
|
|
self.fh.flush()
|
|||
|
|
|
|||
|
|
fh = open(log_path, 'a', buffering=1)
|
|||
|
|
sys.stdout = _Tee(sys.stdout, fh)
|
|||
|
|
sys.stderr = _Tee(sys.stderr, fh)
|
|||
|
|
|
|||
|
|
|
|||
|
|
def _level_of(stem):
|
|||
|
|
"""檔名 stem(無 .nii.gz)的前綴若為 LABEL_MAP 的 level 名則回傳,否則 None。"""
|
|||
|
|
prefix = stem.split('_')[0]
|
|||
|
|
return prefix if prefix in LEVEL_NAMES else None
|
|||
|
|
|
|||
|
|
|
|||
|
|
def detect_source_level(mask_path, vol_dir):
|
|||
|
|
"""由 MASK 相對於 vol_dir 的路徑推 (source, level);推不出時 raise。"""
|
|||
|
|
rel = os.path.relpath(os.path.abspath(mask_path), os.path.abspath(vol_dir))
|
|||
|
|
parts = rel.split(os.sep)
|
|||
|
|
if len(parts) < 1 or not parts[-1].endswith('.nii.gz'):
|
|||
|
|
raise ValueError(f'{mask_path}:需為 .nii.gz')
|
|||
|
|
stem = parts[-1][:-len('.nii.gz')]
|
|||
|
|
level = _level_of(stem)
|
|||
|
|
if level is None:
|
|||
|
|
raise ValueError(f'{mask_path}:檔名前綴推不出 level(預期 <L1..L5/T#>/C#>_...)')
|
|||
|
|
if len(parts) >= 2 and parts[-2] == 'rotated':
|
|||
|
|
return 'rotated', level
|
|||
|
|
suffix = stem.split('_', 1)[1] if '_' in stem else ''
|
|||
|
|
if suffix in SOURCES and suffix != 'rotated':
|
|||
|
|
return suffix, level
|
|||
|
|
raise ValueError(f'{mask_path}:推不出 source(目錄需為 .../rotated/ 或檔名需為 '
|
|||
|
|
f'<level>_{{{"|".join(SOURCES)}}}.nii.gz)')
|
|||
|
|
|
|||
|
|
|
|||
|
|
def discover_levels(vol_dir):
|
|||
|
|
"""volume 目錄(含 crop/、rotated/)出現的所有 level,依 label id 排序。
|
|||
|
|
未旋轉檔:新世代在 crop/ 子資料夾、舊世代在頂層,兩者都掃。"""
|
|||
|
|
lv = set()
|
|||
|
|
for d in (vol_dir, os.path.join(vol_dir, 'crop'), os.path.join(vol_dir, 'rotated')):
|
|||
|
|
if not os.path.isdir(d):
|
|||
|
|
continue
|
|||
|
|
for f in os.listdir(d):
|
|||
|
|
if f.endswith('.nii.gz'):
|
|||
|
|
l = _level_of(f[:-len('.nii.gz')])
|
|||
|
|
if l is not None:
|
|||
|
|
lv.add(l)
|
|||
|
|
return sorted(lv, key=lambda l: _LABEL_ID[l])
|
|||
|
|
|
|||
|
|
|
|||
|
|
def find_native_label(name):
|
|||
|
|
"""LABEL_ROOT 各子目錄找 <name>_seg.nii.gz;回傳 path 或 None。"""
|
|||
|
|
for sub in LABEL_SUBDIRS:
|
|||
|
|
p = os.path.join(LABEL_ROOT, sub, f'{name}_seg.nii.gz')
|
|||
|
|
if os.path.isfile(p):
|
|||
|
|
return p
|
|||
|
|
return None
|
|||
|
|
|
|||
|
|
|
|||
|
|
def _cc_bbox05(mask_img):
|
|||
|
|
"""最大 26-CC bbox:mask_img(0/1)-> [x0,y0,z0,xs,ys,zs]。"""
|
|||
|
|
cc = _largest_cc_bbox(mask_img)
|
|||
|
|
if cc is None:
|
|||
|
|
raise RuntimeError('empty level mask after largest-CC extraction')
|
|||
|
|
return cc[1]
|
|||
|
|
|
|||
|
|
|
|||
|
|
def _disk_box(orig_img, disk_img, flips, spacing):
|
|||
|
|
"""由 disk 檔 origin 回復該檔的裁切 box [x0,y0,z0,xs,ys,zs](x,y,z 序)。
|
|||
|
|
|
|||
|
|
standardize_affine 翻軸的檔:disk origin (LPS) = -pre-standardize
|
|||
|
|
origin(資料被鏡射到世界原點鏡射面);未翻軸:disk origin =
|
|||
|
|
pre-standardize origin = O_full + box0·(d·h)。spacing 為該檔每軸 mm
|
|||
|
|
(0.5mm 檔 = [0.5,0.5,0.5];原解析度檔 = 原 CT spacing)。
|
|||
|
|
非整數(>1e-3)時 raise(幾何模型不符)。"""
|
|||
|
|
o_full = np.array(orig_img.GetOrigin(), dtype=float)
|
|||
|
|
o_disk = np.array(disk_img.GetOrigin(), dtype=float)
|
|||
|
|
d = np.diag(np.array(orig_img.GetDirection(), dtype=float).reshape(3, 3))
|
|||
|
|
s = np.asarray(spacing, dtype=float)
|
|||
|
|
if s.ndim == 0:
|
|||
|
|
s = np.array([float(s)] * 3)
|
|||
|
|
n = np.array(disk_img.GetSize(), dtype=float)
|
|||
|
|
x0 = [((-o_disk[i] - o_full[i]) if i in flips else (o_disk[i] - o_full[i]))
|
|||
|
|
/ (d[i] * s[i]) for i in range(3)]
|
|||
|
|
for i, v in enumerate(x0):
|
|||
|
|
if abs(v - round(v)) > 1e-3:
|
|||
|
|
raise ValueError(f'{disk_img} 翻軸軸 {i} 隱含 box 非整數 ({v:.4f}),'
|
|||
|
|
f'幾何模型不符')
|
|||
|
|
return [int(round(v)) for v in x0] + [int(v) for v in n]
|
|||
|
|
|
|||
|
|
|
|||
|
|
def _recompute_box05(lb_arr, lab_id, lb_img, ct_img):
|
|||
|
|
"""fallback:native label 重算(未翻轉慣例)的最大 26-CC 0.5mm bbox
|
|||
|
|
(與 seg_bone 的 bbox2 同定義)。舊世代 ap_flip=True 個案的 disk 裁切
|
|||
|
|
位置與此可能有 ≤1 voxel 差(舊版 flip 路徑),故優先用 _disk_box。"""
|
|||
|
|
m0_img = sitk.GetImageFromArray((lb_arr == lab_id).astype(np.uint8))
|
|||
|
|
m0_img.CopyInformation(lb_img)
|
|||
|
|
bin_lin = resample_img(sitk.Cast(m0_img, sitk.sitkFloat32))
|
|||
|
|
m_full = (sitk.GetArrayFromImage(bin_lin) > 0.5).astype(np.uint8)
|
|||
|
|
mf_img = sitk.GetImageFromArray(m_full)
|
|||
|
|
mf_img.CopyInformation(bin_lin)
|
|||
|
|
return [int(v) for v in _cc_bbox05(mf_img)]
|
|||
|
|
|
|||
|
|
|
|||
|
|
def rebuild_level_boxes(vol_dir, level, lab_id, lb_img, lb_arr, orig_img,
|
|||
|
|
flips):
|
|||
|
|
"""回復 level 各檔的裁切 box:優先由 disk 檔 origin 回復(精確、含舊世代
|
|||
|
|
ap_flip 差異);該檔缺失時 fallback native label 重算(未翻轉慣例)並
|
|||
|
|
警告。0.5mm 檔(smd_resampled/binary_sdf/roi 共用 box2、binary_nn 自己
|
|||
|
|
的 nn box)+ 原解析度檔(binary、smd)都在此。"""
|
|||
|
|
name = os.path.basename(os.path.abspath(vol_dir))
|
|||
|
|
boxes = {}
|
|||
|
|
shared = None
|
|||
|
|
for sname in ('smd_resampled', 'binary_sdf', 'roi'):
|
|||
|
|
p = level_file_path(vol_dir, level, sname)
|
|||
|
|
if os.path.exists(p):
|
|||
|
|
try:
|
|||
|
|
b = _disk_box(orig_img, sitk.ReadImage(p), flips, 0.5)
|
|||
|
|
if shared is None:
|
|||
|
|
shared = b
|
|||
|
|
elif b != shared and max(abs(a - c) for a, c in zip(b, shared)) > 0:
|
|||
|
|
logger.warning(f'{name} {level}: {sname} disk box '
|
|||
|
|
f'{b} != smd_resampled {shared}(應相同)')
|
|||
|
|
except ValueError as e:
|
|||
|
|
logger.warning(f'{name} {level}: {sname} disk box 回復失敗'
|
|||
|
|
f'({e})')
|
|||
|
|
if shared is None:
|
|||
|
|
shared = _recompute_box05(lb_arr, lab_id, lb_img, orig_img)
|
|||
|
|
shared = list(shared)
|
|||
|
|
logger.warning(f'{name} {level}: 無 0.5mm 檔案可回復 disk box,'
|
|||
|
|
f'fallback native label 重算 {shared}(ap_flip=True 舊'
|
|||
|
|
f'世代時可能 ≤1 voxel 偏差)')
|
|||
|
|
boxes['smd_resampled'] = list(shared)
|
|||
|
|
boxes['binary_sdf'] = list(shared)
|
|||
|
|
boxes['roi'] = list(shared)
|
|||
|
|
|
|||
|
|
nn_path = level_file_path(vol_dir, level, 'binary_nn')
|
|||
|
|
if os.path.exists(nn_path):
|
|||
|
|
boxes['binary_nn'] = _disk_box(orig_img, sitk.ReadImage(nn_path),
|
|||
|
|
flips, 0.5)
|
|||
|
|
else:
|
|||
|
|
nn_img = resample_img(lb_img, is_label=True)
|
|||
|
|
mn = (sitk.GetArrayFromImage(nn_img) == lab_id).astype(np.uint8)
|
|||
|
|
mn_img = sitk.GetImageFromArray(mn)
|
|||
|
|
mn_img.CopyInformation(nn_img)
|
|||
|
|
nn_res = _largest_cc_bbox(mn_img)
|
|||
|
|
if nn_res is not None:
|
|||
|
|
boxes['binary_nn'] = [int(v) for v in nn_res[1]]
|
|||
|
|
else:
|
|||
|
|
logger.warning(f'{name} {level}: NN fallback 無連通區域')
|
|||
|
|
|
|||
|
|
for sname, fb in (('binary', 'm0'), ('smd', 'smd4')):
|
|||
|
|
p = level_file_path(vol_dir, level, sname)
|
|||
|
|
try:
|
|||
|
|
boxes[sname] = _disk_box(orig_img, sitk.ReadImage(p), flips,
|
|||
|
|
orig_img.GetSpacing())
|
|||
|
|
except (ValueError, OSError) as e:
|
|||
|
|
m0_img = sitk.GetImageFromArray((lb_arr == lab_id).astype(np.uint8))
|
|||
|
|
m0_img.CopyInformation(lb_img)
|
|||
|
|
bbox_orig = [int(v) for v in _cc_bbox05(m0_img)]
|
|||
|
|
boxes[sname] = bbox_orig if fb == 'm0' else \
|
|||
|
|
margined_box(bbox_orig, orig_img.GetSize(), margin=4)
|
|||
|
|
logger.warning(f'{name} {level}: {sname} disk box 回復失敗'
|
|||
|
|
f'({e}),fallback native label')
|
|||
|
|
return boxes
|
|||
|
|
|
|||
|
|
|
|||
|
|
def rebuild_rotated_section(vol_dir, level):
|
|||
|
|
"""重算 rotated 座標:template 選擇與 R/center 重算(與 _write_rotated_level
|
|||
|
|
同路徑,確定性)、fstart 由 rotated/target 檔的 disk 幾何回復,並做正向
|
|||
|
|
identity 驗證(mismatch > 1e-4 raise)。回傳 section dict 或
|
|||
|
|
(None, reason)。"""
|
|||
|
|
tpl = _load_unrotated_template(vol_dir, level)
|
|||
|
|
if tpl is None:
|
|||
|
|
return None, 'no unrotated template'
|
|||
|
|
t_img, t_arr, is_smd = tpl
|
|||
|
|
rot_dir = os.path.join(vol_dir, 'rotated')
|
|||
|
|
rot_path = None
|
|||
|
|
for s in ('smd_resampled', 'binary_sdf', 'label'):
|
|||
|
|
p = os.path.join(rot_dir, f'{level}_{s}.nii.gz')
|
|||
|
|
if os.path.exists(p):
|
|||
|
|
rot_path = p
|
|||
|
|
break
|
|||
|
|
if rot_path is None:
|
|||
|
|
return None, 'no rotated file'
|
|||
|
|
rot_img = sitk.ReadImage(rot_path)
|
|||
|
|
fstart = _fstart_from_geometries(rot_img, t_img)
|
|||
|
|
if fstart is None:
|
|||
|
|
return None, 'rotated/unrotated geometry mismatch'
|
|||
|
|
bin_arr = _bone_mask(t_arr, is_smd).astype(np.uint8)
|
|||
|
|
sym = best_symmetry_plane(bin_arr)
|
|||
|
|
symp = best_upper_endplate_plane(bin_arr)
|
|||
|
|
R, c = compute_normalizing_rotation(bin_arr, sym, symp)
|
|||
|
|
fwd = _forward_identity_check(vol_dir, level, R, c, t_img, t_arr, is_smd)
|
|||
|
|
if fwd is not None and fwd > 1e-4:
|
|||
|
|
raise RuntimeError(f'{level}: 正向重旋轉 mismatch {fwd:.3e} > 1e-4'
|
|||
|
|
f'(重算 R/center 與生成參數不符;拒絕寫入)')
|
|||
|
|
if fwd is None:
|
|||
|
|
logger.warning(f'{level}: 無可比對的 rotated 檔,R/center 未經驗證')
|
|||
|
|
return {
|
|||
|
|
'template': 'smd_resampled' if is_smd else 'binary_nn',
|
|||
|
|
'R': [[float(v) for v in row] for row in R],
|
|||
|
|
'center': [float(v) for v in c],
|
|||
|
|
'start': [int(v) for v in fstart],
|
|||
|
|
'size': [int(v) for v in rot_img.GetSize()],
|
|||
|
|
}, ('forward_mismatch=%.1e' % fwd) if fwd is not None else 'unverified'
|
|||
|
|
|
|||
|
|
|
|||
|
|
def _detect_ap_flip(vol_dir, meta, lb_arr, ap_db, max_levels=3):
|
|||
|
|
"""disk 資料是否含 ap_flip 的對位檢測:以「template 骨頭 -> 原始 grid」
|
|||
|
|
與 native label 的 bone IoU 比對兩種假設,取高者。
|
|||
|
|
|
|||
|
|
metadata db 的 ap_flip 只反映【新 pipeline 世代】(09-07 起)的行為;
|
|||
|
|
舊世代(如 standardized-xfr-2,09-06 產出)disk 上沒有 y 翻轉,db 值
|
|||
|
|
(可能為 True)不可直接用於舊世代資料 —— 故重建時一律以對位結果為準,
|
|||
|
|
db 值僅作為先驗(試驗順序 / 高 IoU 時短路)。"""
|
|||
|
|
import copy
|
|||
|
|
lvls = [l for l in meta['levels'] if 'rotated' in meta['levels'][l]]
|
|||
|
|
name = os.path.basename(os.path.abspath(vol_dir))
|
|||
|
|
if not lvls:
|
|||
|
|
chosen = bool(ap_db) if ap_db is not None else False
|
|||
|
|
logger.warning(f'{name}: 無 rotated level 可驗證 ap_flip;'
|
|||
|
|
f'{"用 db 值" if ap_db is not None else "缺 db 值"} {chosen}')
|
|||
|
|
return chosen
|
|||
|
|
order = [bool(ap_db), not bool(ap_db)] if ap_db is not None else [False, True]
|
|||
|
|
best_hyp, best_iou = None, -1.0
|
|||
|
|
for hyp in order:
|
|||
|
|
m = copy.deepcopy(meta)
|
|||
|
|
m['ap_flip'] = hyp
|
|||
|
|
ious = []
|
|||
|
|
for level in lvls[:max_levels]:
|
|||
|
|
tpl = m['levels'][level]['rotated']['template']
|
|||
|
|
a = sitk.GetArrayFromImage(
|
|||
|
|
sitk.ReadImage(level_file_path(vol_dir, level, tpl)))
|
|||
|
|
bone = ((a < 0.5) if tpl == 'smd_resampled' else (a > 0)).astype(np.uint8)
|
|||
|
|
if not bone.any():
|
|||
|
|
continue
|
|||
|
|
out = mask_to_original(bone, m, level, tpl)
|
|||
|
|
lab = (lb_arr == _LABEL_ID[level])
|
|||
|
|
on = out > 0
|
|||
|
|
union = int((on | lab).sum())
|
|||
|
|
ious.append(int((on & lab).sum()) / union if union else 0.0)
|
|||
|
|
iou = float(np.mean(ious)) if ious else 0.0
|
|||
|
|
logger.info(f'{name}: ap_flip={hyp} -> 骨頭 native IoU={iou:.4f}')
|
|||
|
|
if iou > best_iou:
|
|||
|
|
best_iou, best_hyp = iou, hyp
|
|||
|
|
if ap_db is not None and hyp == bool(ap_db) and iou >= 0.95:
|
|||
|
|
break # 先驗假設已充分對位,免試另一半
|
|||
|
|
if ap_db is not None and best_hyp != bool(ap_db):
|
|||
|
|
logger.warning(f'{name}: ap_flip 對位檢測={best_hyp} 與 metadata db='
|
|||
|
|
f'{bool(ap_db)} 不同(db 值屬新 pipeline 世代;disk 資料'
|
|||
|
|
f'應為另一世代),採對位檢測結果')
|
|||
|
|
if best_iou < 0.9:
|
|||
|
|
logger.warning(f'{name}: ap_flip 對位檢測最佳 IoU={best_iou:.4f} < 0.9,'
|
|||
|
|
f'座標鏈可能有誤,請人工核對')
|
|||
|
|
return best_hyp
|
|||
|
|
|
|||
|
|
|
|||
|
|
def rebuild_transform(vol_dir, orig_path=None):
|
|||
|
|
"""舊 volume(無 transform.json):從 disk 輸出 + data_root 原始 CT +
|
|||
|
|
native label + metadata db 重建完整 transform.json 並寫回。
|
|||
|
|
回傳 meta。所有推導與 pipeline 記錄同式(box 由 disk 檔 origin 回復、
|
|||
|
|
缺檔 fallback native label 重算、flip 軸由原 CT direction 推斷、
|
|||
|
|
rotated R/center 重算 + 正向 identity 驗證、ap_flip 以對位 IoU 檢測
|
|||
|
|
(db 值僅先驗,見 _detect_ap_flip))。"""
|
|||
|
|
name = os.path.basename(os.path.abspath(vol_dir))
|
|||
|
|
if orig_path is None:
|
|||
|
|
orig_path, _ = find_original_ct(name)
|
|||
|
|
if orig_path is None:
|
|||
|
|
raise FileNotFoundError(f'{DATA_ROOT}*/{name}.nii.gz 找不到')
|
|||
|
|
orig_img = sitk.ReadImage(orig_path)
|
|||
|
|
flip_axes = std_flip_axes_for_direction(orig_img.GetDirection())
|
|||
|
|
lb_path = find_native_label(name)
|
|||
|
|
if lb_path is None:
|
|||
|
|
raise FileNotFoundError(f'{LABEL_ROOT}*/{name}_seg.nii.gz 找不到')
|
|||
|
|
lb_img = sitk.ReadImage(lb_path)
|
|||
|
|
lb_arr = sitk.GetArrayFromImage(lb_img)
|
|||
|
|
|
|||
|
|
levels = discover_levels(vol_dir)
|
|||
|
|
if not levels:
|
|||
|
|
raise ValueError(f'{vol_dir}: 找不到任何 <level>_*.nii.gz')
|
|||
|
|
level_entries = {}
|
|||
|
|
for level in levels:
|
|||
|
|
boxes = rebuild_level_boxes(vol_dir, level, _LABEL_ID[level], lb_img,
|
|||
|
|
lb_arr, orig_img, flip_axes)
|
|||
|
|
lv = {'label': _LABEL_ID[level], 'std_flip_axes': flip_axes,
|
|||
|
|
'boxes': boxes}
|
|||
|
|
sec, note = rebuild_rotated_section(vol_dir, level)
|
|||
|
|
if sec is not None:
|
|||
|
|
lv['rotated'] = sec
|
|||
|
|
logger.info(f'{name} {level}: rebuilt (rotated verified: {note})')
|
|||
|
|
else:
|
|||
|
|
logger.info(f'{name} {level}: rebuilt (no rotated section: {note})')
|
|||
|
|
level_entries[level] = lv
|
|||
|
|
orig_geom = img_geom(orig_img)
|
|||
|
|
meta = build_volume_meta(name, orig_path, orig_geom,
|
|||
|
|
resampled05_geom_from_original(orig_geom),
|
|||
|
|
False, level_entries)
|
|||
|
|
ap_db = load_ap_flip().get(name)
|
|||
|
|
meta['ap_flip'] = _detect_ap_flip(vol_dir, meta, lb_arr, ap_db)
|
|||
|
|
save_transform(vol_dir, meta)
|
|||
|
|
logger.info(f'{name}: rebuild -> {transform_path(vol_dir)} '
|
|||
|
|
f'({len(level_entries)} level(s), ap_flip={meta["ap_flip"]}, '
|
|||
|
|
f'std_flip_axes={flip_axes})')
|
|||
|
|
return meta
|
|||
|
|
|
|||
|
|
|
|||
|
|
def _check_orig_geometry(meta, orig_img, src='recorded'):
|
|||
|
|
rec = meta['original']
|
|||
|
|
if (list(orig_img.GetSize()) != rec['size']
|
|||
|
|
or any(abs(float(a) - b) > 1e-6
|
|||
|
|
for a, b in zip(orig_img.GetSpacing(), rec['spacing']))):
|
|||
|
|
sys.exit(f'原始 CT 幾何與記錄不符({src}):size {orig_img.GetSize()} '
|
|||
|
|
f'spacing {orig_img.GetSpacing()} vs 記錄 {rec["size"]} '
|
|||
|
|
f'{rec["spacing"]};座標鏈以記錄值為準,無法映射')
|
|||
|
|
|
|||
|
|
|
|||
|
|
def main():
|
|||
|
|
parser = argparse.ArgumentParser(
|
|||
|
|
description='Inverse-map a standardized-grid segmentation mask to the '
|
|||
|
|
'original CT coordinates (data_root) via transform.json.')
|
|||
|
|
parser.add_argument('vol_dir', help='standardized volume dir (含 transform.json)')
|
|||
|
|
parser.add_argument('mask', help='mask nifti(在 source grid 上)')
|
|||
|
|
parser.add_argument('--source', default=None,
|
|||
|
|
choices=['auto'] + list(SOURCES),
|
|||
|
|
help='source grid(default: 由檔名自動判定)')
|
|||
|
|
parser.add_argument('--level', default=None,
|
|||
|
|
help='level(default: 由檔名前綴自動判定)')
|
|||
|
|
parser.add_argument('--out', default=None,
|
|||
|
|
help='輸出 nifti(default: <mask>_to_original.nii.gz)')
|
|||
|
|
parser.add_argument('--orig', default=None,
|
|||
|
|
help='原始 CT path(default: transform.json 記錄值)')
|
|||
|
|
parser.add_argument('--rebuild', action='store_true',
|
|||
|
|
help='transform.json 缺失時重建(舊世代 volume)')
|
|||
|
|
parser.add_argument('--chunk', type=int, default=16,
|
|||
|
|
help='反向映射的 z 區塊 slice 數(記憶體控制, default 16)')
|
|||
|
|
args = parser.parse_args()
|
|||
|
|
|
|||
|
|
if not os.path.isdir(args.vol_dir):
|
|||
|
|
sys.exit(f'vol_dir 不存在:{args.vol_dir}')
|
|||
|
|
if not os.path.isfile(args.mask):
|
|||
|
|
sys.exit(f'mask 不存在:{args.mask}')
|
|||
|
|
|
|||
|
|
os.makedirs(LOG_DIR, exist_ok=True)
|
|||
|
|
log_path = os.path.join(LOG_DIR,
|
|||
|
|
f'xfr_inverse_transform_{time.strftime("%Y%m%d_%H%M%S")}.log')
|
|||
|
|
setup_tee(log_path)
|
|||
|
|
logging.basicConfig(level=logging.INFO,
|
|||
|
|
format='%(asctime)s [%(levelname)s] %(name)s: %(message)s',
|
|||
|
|
datefmt='%Y-%m-%d %H:%M:%S')
|
|||
|
|
logger.info(f'Log file: {log_path}')
|
|||
|
|
logger.info(f'Command: {sys.executable} {" ".join(sys.argv)}')
|
|||
|
|
|
|||
|
|
tp = transform_path(args.vol_dir)
|
|||
|
|
if os.path.exists(tp):
|
|||
|
|
meta = load_transform(args.vol_dir)
|
|||
|
|
elif args.rebuild:
|
|||
|
|
meta = rebuild_transform(args.vol_dir, args.orig)
|
|||
|
|
else:
|
|||
|
|
sys.exit(f'{tp} 不存在;重跑 pipeline 或在既有 volume 上加 --rebuild')
|
|||
|
|
|
|||
|
|
source = args.source if (args.source and args.source != 'auto') else None
|
|||
|
|
level = args.level
|
|||
|
|
if source is None or level is None:
|
|||
|
|
a_source, a_level = detect_source_level(args.mask, args.vol_dir)
|
|||
|
|
source = source if source is not None else a_source
|
|||
|
|
level = level if level is not None else a_level
|
|||
|
|
|
|||
|
|
name = meta.get('name') or os.path.basename(os.path.abspath(args.vol_dir))
|
|||
|
|
if level not in meta['levels']:
|
|||
|
|
sys.exit(f'{level} 不在 transform.json({sorted(meta["levels"])})')
|
|||
|
|
if source == 'rotated' and 'rotated' not in meta['levels'][level]:
|
|||
|
|
sys.exit(f'{level} 無 rotated 記錄(未跑 _write_rotated_level 的 level '
|
|||
|
|
f'沒有 rotated grid 座標鏈)')
|
|||
|
|
if source in ('smd_resampled', 'binary_sdf', 'roi', 'binary_nn',
|
|||
|
|
'binary', 'smd') and source not in meta['levels'][level]['boxes']:
|
|||
|
|
sys.exit(f'{level} 無 {source} 記錄(該檔未產出)')
|
|||
|
|
|
|||
|
|
if args.orig:
|
|||
|
|
orig_img = sitk.ReadImage(args.orig)
|
|||
|
|
_check_orig_geometry(meta, orig_img, '--orig')
|
|||
|
|
else:
|
|||
|
|
p = meta['original']['path']
|
|||
|
|
if not os.path.exists(p):
|
|||
|
|
sys.exit(f'記錄的原始 CT 不存在:{p}(可用 --orig 指定)')
|
|||
|
|
orig_img = sitk.ReadImage(p)
|
|||
|
|
|
|||
|
|
mask_img = sitk.ReadImage(args.mask)
|
|||
|
|
t0 = time.time()
|
|||
|
|
arr = mask_to_original(mask_img, meta, level, source, chunk=args.chunk)
|
|||
|
|
dt = time.time() - t0
|
|||
|
|
|
|||
|
|
n_on = int((arr != 0).sum())
|
|||
|
|
ct = sitk.GetArrayFromImage(orig_img)
|
|||
|
|
if n_on:
|
|||
|
|
hu = ct[arr != 0]
|
|||
|
|
frac_body = float((hu > -100).mean())
|
|||
|
|
med_hu = float(np.median(hu))
|
|||
|
|
logger.info(f'alignment (info): nonzero={n_on}, in-body(HU>-100)='
|
|||
|
|
f'{frac_body:.3f}, median HU={med_hu:.0f}')
|
|||
|
|
else:
|
|||
|
|
logger.warning('反向映射結果全零(mask 可能全空或 grid 錯位)')
|
|||
|
|
|
|||
|
|
out_path = args.out or (args.mask[:-len('.nii.gz')] + '_to_original.nii.gz')
|
|||
|
|
from imaging.transforms import save_as_original
|
|||
|
|
save_as_original(arr, orig_img, out_path)
|
|||
|
|
logger.info(f'saved {out_path} ({name} {level} {source} -> original, '
|
|||
|
|
f'nonzero={n_on}, {dt:.1f}s)')
|
|||
|
|
|
|||
|
|
|
|||
|
|
if __name__ == '__main__':
|
|||
|
|
main()
|