CBT_project/xfr_inverse_transform.py

507 lines
23 KiB
Python
Raw Permalink Normal View History

#!/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 Rcenter
start index 空間映射不依賴 standardized 輸出的物理 headerNN
order=0輸出幾何 = 原始 CTpixel type = 輸入 mask
支援的 source grid--source預設由檔名自動判定
rotated <vol>/rotated/<L>_*.nii.gz 上的任何 mask label map
自己跑的分割結果整格同幾何
smd_resampled <vol>[/crop]/<L>_smd_resampled.nii.gz0.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 檔一致 gridlevel 取檔名前綴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 niftilabel / 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 bboxmask_img0/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):
"""fallbacknative 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 共用 box2binary_nn 自己
nn box+ 原解析度檔binarysmd都在此"""
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-209-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 griddefault: 由檔名自動判定)')
parser.add_argument('--level', default=None,
help='leveldefault: 由檔名前綴自動判定)')
parser.add_argument('--out', default=None,
help='輸出 niftidefault: <mask>_to_original.nii.gz')
parser.add_argument('--orig', default=None,
help='原始 CT pathdefault: 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()