#!/home/xfr/.conda/envs/cbt/bin/python """把標準化 grid 上的 segmentation mask 反向映射回 data_root 原始 CT grid。 座標參數來自 xfr_preprocess 記錄的 /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 /rotated/_*.nii.gz 上的任何 mask(含 label map、 自己跑的分割結果;整格同幾何) smd_resampled [/crop]/_smd_resampled.nii.gz(0.5mm 裁切) binary_sdf [/crop]/_binary_sdf.nii.gz roi [/crop]/_roi.nii.gz binary_nn [/crop]/_binary_nn.nii.gz binary [/crop]/_binary.nii.gz(原解析度) smd [/crop]/_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 等)用這個。 輸出預設:_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(預期 /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'_{{{"|".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 各子目錄找 _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}: 找不到任何 _*.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: _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()