Introduce a robust coordinate transformation system to manage the relationship between original CT space and rotated/standardized segmentation spaces. This includes a new directory hierarchy to separate unrotated crops from rotated outputs and utility functions for geometric mapping. Key changes: - Implement `imaging/transforms.py` to handle bounding box metadata, affine standardization, and coordinate mapping between spaces. - Restructure dataset output: unrotated segmentation files (binary, SDF, ROI, etc.) are now stored in a `<vol>/crop/` subdirectory to distinguish them from `<vol>/rotated/` aligned versions. - Add `level_file_path` utility to abstract file discovery across legacy (top-level) and new (crop-based) directory structures. - Enhance `seg_bone` to capture and export bounding box metadata (`bbox2`, `nn_bbox`, `bbox_orig`) into `transform.json`. - Implement `xfr_cbt_native.py` for mapping screw positions back to original CT space. - Update preprocessing and visualization scripts to support the new directory layout and transformation metadata. - Improve TinyDB metadata migration logic to prevent accidental corruption of existing database structures.
507 lines
No EOL
23 KiB
Python
507 lines
No EOL
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() |