CBT_project/imaging/transforms.py
Xiao Furen d167c1f7c7 feat(imaging): implement coordinate transformation pipeline and directory restructuring
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.
2026-09-09 13:39:47 +08:00

304 lines
No EOL
12 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""Per-level 座標變換記錄transform.json與反向映射把標準化
0.5mm / rotatedgrid 上的 segmentation mask 映回 data_root 原始 CT grid。
座標鏈(純 index 空間、(x,y,z) 向量慣例、array = (z,y,x);不依賴 standardized
輸出的物理 header —— standardize_affine 的 origin 處理不可靠2026-09-08 已
驗證xfr_cbt_native / xfr_orig_labels 同結論):
original index odata_root CT
-> 真 0.5mm full grid: t = o * (so / 0.5)
-> 工作系ap_flip 後): ap_flip 時 g_y = N05_y - 1 - t_y
-> 未旋轉檔 pre-standardize: p = g - boxbox = 該檔在 full grid 的裁切)
-> diskstandardize 後): std_flip_axes 的軸: disk_i = N_i - 1 - p_i
-> rotated disk: a = R (p_tpl - center) + center - start
transform.json 記錄每一步所需的參數(寫檔時即為 ground truth反向映射不需
重新估算任何平面 / 幾何):
{
"version": 1,
"name": "<volume name>",
"original": {"path", "size", "spacing", "direction", "origin"},
"resampled05": {"size", "spacing", "direction", "origin"}, // 0.5mm full grid
"ap_flip": false,
"levels": {
"L1": {
"label": 20,
"std_flip_axes": [0, 1], // standardize_affine 實際翻的軸
"boxes": { // [x0,y0,z0,xs,ys,zs]x,y,z 序);
"smd_resampled": [...], // 0.5mm 檔: 0.5mm full gridap_flip 後)
"binary_sdf": [...], // 原解析度檔 (binary/smd): 原 index
"roi": [...], // ap_flip 後)
"binary_nn": [...],
"binary": [...],
"smd": [...]
},
"rotated": { // _write_rotated_level 補寫
"template": "smd_resampled", // 或 "binary_nn"fallback
"R": [[..]], // 作用於 (x,y,z)forward dest = R(src-c)+c
"center": [cx, cy, cz], // 未旋轉檔 diskstandardize 後index 系
"start": [sx, sy, sz], // rotated 檔 origin模板 index 系)
"size": [nx, ny, nz]
}
}
}
}
original_to_source 回傳 (M (3,3), t (3,))source disk index = M @ original
index + t連續座標x,y,z 系。mask_to_original 對每個 original voxel 以
order=0NN在 source mask 上採樣,得原始 grid 的 mask。
"""
import json
import os
import numpy as np
import SimpleITK as sitk
from scipy.ndimage import map_coordinates
import logging
logger = logging.getLogger('imaging.transforms')
TRANSFORM_FILENAME = 'transform.json'
# 0.5mm 標準化 grid 的 sourceboxes 記在 0.5mm full grid、ap_flip 後)
SOURCES_05 = ('smd_resampled', 'binary_sdf', 'roi', 'binary_nn')
# 原解析度 grid 的 sourceboxes 記在原始 index、ap_flip 後)
SOURCES_ORIG = ('binary', 'smd')
SOURCES = SOURCES_05 + SOURCES_ORIG + ('rotated',)
def transform_path(vol_dir):
return os.path.join(vol_dir, TRANSFORM_FILENAME)
def level_file_path(vol_dir, level, source):
"""<level>_<source>.nii.gz 的位置(未旋轉、裁切 level bbox 的檔):
新世代xfr-3 起)落在 <vol>/crop/ 子資料夾(旋轉版在 <vol>/rotated/
舊世代xfr-2 等)落在 <vol>/ 頂層。回傳先存在者crop/ 優先);
兩者皆無時回 crop/ 路徑(呼叫端以 os.path.exists 判定)."""
p_crop = os.path.join(vol_dir, 'crop', f'{level}_{source}.nii.gz')
if os.path.exists(p_crop):
return p_crop
p_top = os.path.join(vol_dir, f'{level}_{source}.nii.gz')
if os.path.exists(p_top):
return p_top
return p_crop
def load_transform(vol_dir):
p = transform_path(vol_dir)
if not os.path.exists(p):
raise FileNotFoundError(
f'{p} 不存在(需先跑 xfr_preprocess pipeline'
f'或 xfr_inverse_transform.py --rebuild')
with open(p) as f:
return json.load(f)
def save_transform(vol_dir, meta):
p = transform_path(vol_dir)
tmp = f'{p}.tmp'
with open(tmp, 'w') as f:
json.dump(meta, f, indent=2, ensure_ascii=False)
os.replace(tmp, p)
def merge_rotated_into_transform(vol_dir, level, sec):
"""讀入既有 transform.json補 / 覆寫 levels[level]['rotated'] 後寫回
_write_rotated_level 用base 部分由 process_single_image 寫入)。"""
meta = load_transform(vol_dir)
lv = meta['levels'].setdefault(level, {})
lv['rotated'] = sec
save_transform(vol_dir, meta)
def img_geom(img):
"""sitk img -> {"size"(x,y,z), "spacing", "direction"(row-major 9), "origin"}。"""
return {
'size': [int(v) for v in img.GetSize()],
'spacing': [float(v) for v in img.GetSpacing()],
'direction': [float(v) for v in img.GetDirection()],
'origin': [float(v) for v in img.GetOrigin()],
}
def require_diagonal(direction, what='direction'):
d = np.asarray(direction, dtype=float).reshape(3, 3)
diag = np.diag(d)
if not (np.allclose(d, np.diag(diag), atol=1e-6)
and np.allclose(np.abs(diag), 1.0, atol=1e-6)):
raise ValueError(f'{what} 非對角 ±1本座標鏈只支援對角 direction: '
f'{d.tolist()}')
return diag
def std_flip_axes_for_direction(direction):
"""standardize_affine 會翻的軸:儲存 NIfTI affineRAS對角 < 0 者。
RAS = LPS 的 x/y 取反ras_diag = (-d_x, -d_y, +d_z)。
direction 非對角 ±1 時 raise ValueError。"""
d = require_diagonal(direction, 'direction')
ras = np.array([-d[0], -d[1], d[2]])
return [int(i) for i in range(3) if ras[i] < 0]
def resampled05_geom_from_original(orig_geom):
"""resample_img 的 0.5mm full grid 幾何origin/direction 同原、尺寸 ceil"""
n = np.asarray(orig_geom['size'], dtype=float)
s = np.asarray(orig_geom['spacing'], dtype=float)
return {
'size': [max(1, int(vv)) for vv in np.ceil(n * s / 0.5 - 1e-6)],
'spacing': [0.5, 0.5, 0.5],
'direction': [float(v) for v in orig_geom['direction']],
'origin': [float(v) for v in orig_geom['origin']],
}
def margined_box(box, n_xyz, margin=4):
"""seg_bone 的 _bbox_roi 同式box [x0,y0,z0,xs,ys,zs] 對稱外擴 margin、
clamp 到 n_xyzx,y,z 序)。"""
n = np.asarray(n_xyz, dtype=float)
idx = [max(0, int(box[i]) - margin) for i in range(3)]
size = [min(int(n[i]) - idx[i], int(box[i + 3]) + 2 * margin) for i in range(3)]
return idx + size
def build_volume_meta(name, orig_path, orig_geom, r05_geom, ap_flip, level_entries):
return {
'version': 1,
'name': name,
'original': {'path': orig_path, **orig_geom},
'resampled05': r05_geom,
'ap_flip': bool(ap_flip),
'levels': level_entries,
}
# ---------------------------------------------------------------- 座標鏈
def _compose(f1, f2):
"""f1(f2(x)):先 f2 後 f1。f = (M, t)x -> M @ x + t。"""
M1, t1 = f1
M2, t2 = f2
return M1 @ M2, M1 @ t2 + t1
def original_to_source(meta, level, source):
"""原 CT index (x,y,z) -> source disk index (x,y,z) 的連續仿射
(M (3,3), t (3,))s = M @ o + t NN 前座標)。
source ∈ SOURCES。rotated 時 chain 末端加 R 旋轉 + start 平移
template = 未旋轉 0.5mm 檔flip / box 用 template 的)。
所有參數來自 metatransform.json 記錄值,非從 header 估算)。"""
if source not in SOURCES:
raise ValueError(f'unknown source {source!r}; 預期 {SOURCES}')
lv = meta['levels'].get(level)
if lv is None:
raise ValueError(f'level {level!r} 不在 transform.json '
f'({sorted(meta["levels"])})')
if source == 'rotated' and 'rotated' not in lv:
raise ValueError(f'{level} 沒有 rotated 記錄(该 level 未走 _write_rotated_level')
so = np.asarray(meta['original']['spacing'], dtype=float)
n_orig = np.asarray(meta['original']['size'], dtype=float)
n05 = np.asarray(meta['resampled05']['size'], dtype=float)
ap_flip = bool(meta.get('ap_flip', False))
if source == 'rotated':
rot = lv['rotated']
tpl = rot['template']
Mrot = np.asarray(rot['R'], dtype=float)
c = np.asarray(rot['center'], dtype=float)
start = np.asarray(rot['start'], dtype=float)
else:
tpl = source
Mrot = c = start = None
if tpl not in lv['boxes']:
raise ValueError(f'{level} 缺 boxes[{tpl!r}]transform.json 未完)')
box = np.asarray(lv['boxes'][tpl], dtype=float)
res = '05' if tpl in SOURCES_05 else 'orig'
flips = [int(i) for i in lv.get('std_flip_axes', [])]
n_file = box[3:6] # 該檔尺寸(= flip 時的 N
f = (np.eye(3), np.zeros(3))
# 1) original index -> 真 0.5mm full grid index0.5mm source 才需要)
if res == '05':
f = _compose((np.diag(so / 0.5), np.zeros(3)), f)
# 2) -> 工作系ap_flip 後的 0.5mm / 原解析度 full grid index
if ap_flip:
extent = n05[1] if res == '05' else n_orig[1]
M = np.eye(3)
M[1, 1] = -1.0
f = _compose((M, np.array([0.0, float(extent) - 1.0, 0.0])), f)
# 3) -> 未旋轉檔 pre-standardize index裁切 box 的局部系)
f = _compose((np.eye(3), -box[:3]), f)
# 4) -> diskstandardize_affine 翻軸後的存檔 index
for i in flips:
if i >= 3:
raise ValueError(f'std_flip_axes 含非法軸 {i}')
M = np.eye(3)
M[i, i] = -1.0
t = np.zeros(3)
t[i] = float(n_file[i]) - 1.0
f = _compose((M, t), f)
# 5) (rotated) 未旋轉 template disk -> rotated disk
# forward: W = R^T (V - c) + c, V = a + start => a = R(W - c) + c - start
if Mrot is not None:
f = _compose((Mrot, c - Mrot @ c - start), f)
return f
def expected_source_size(meta, level, source):
"""該 source 檔的 (x,y,z) 尺寸rotated 用 rotated.size其餘用 boxes"""
lv = meta['levels'][level]
if source == 'rotated':
return [int(v) for v in lv['rotated']['size']]
box = lv['boxes'][source]
return [int(box[i + 3]) for i in range(3)]
def mask_to_original(mask_img, meta, level, source, chunk=16):
"""把 source grid 上的 masknifti 或 sitk imageNN 反向映射回原始 CT
grid。回傳 (z,y,x) ndarraydtype = 輸入 dtype
chunk每個 z 區塊的 slice 數(原始 CT z 可能上百;分塊控制峰值記憶體)。
mask 的 size 與該 source 檔不一致時 raisegrid 不對應,映射必錯)。"""
M, t = original_to_source(meta, level, source)
arr = mask_img if isinstance(mask_img, np.ndarray) \
else sitk.GetArrayFromImage(mask_img)
if isinstance(mask_img, sitk.Image):
actual = np.asarray(mask_img.GetSize())
else:
actual = np.array(arr.shape[::-1])
expect = np.asarray(expected_source_size(meta, level, source))
if not np.array_equal(actual, expect):
raise ValueError(
f'mask 尺寸 {tuple(map(int, actual))}{level}/{source} 記錄尺寸 '
f'{tuple(map(int, expect))} 不符grid 不對應,無法映射)')
if arr.size == 0:
raise ValueError('輸入 mask 為空')
nx, ny, nz = (int(v) for v in meta['original']['size'])
out = np.zeros((nz, ny, nx), dtype=np.float64)
for z0 in range(0, nz, int(chunk)):
z1 = min(nz, z0 + int(chunk))
oz, oy, ox = np.indices((z1 - z0, ny, nx))
X = ox.astype(np.float64)
Y = oy.astype(np.float64)
Z = (oz + z0).astype(np.float64)
P = np.stack([X, Y, Z], axis=0) # (3, cz, ny, nx) (x,y,z)
S = np.tensordot(M, P, axes=([1], [0])) + t[:, np.newaxis, np.newaxis,
np.newaxis]
out[z0:z1, ...] = map_coordinates(arr, [S[2], S[1], S[0]],
order=0, cval=0.0, mode='constant')
return out.astype(arr.dtype, copy=False)
def save_as_original(arr, orig_img, out_path):
"""arr (z,y,x) -> nifti幾何取 orig_img保留 arr 自身 pixel type"""
out = sitk.GetImageFromArray(arr)
out.SetSpacing(orig_img.GetSpacing())
out.SetDirection(orig_img.GetDirection())
out.SetOrigin(orig_img.GetOrigin())
sitk.WriteImage(out, out_path)