Implement a more robust scoring mechanism for screw optimization and add functionality to generate synthetic X-ray projections (AP and lateral views) from CT data. Key changes: - core: add `generate_cylinder_butt_torch` to create a mask for the screw entrance (0.25mm) to exempt it from bone-breaching penalties. - core: update `cl_score_torch_xfr` to include a diameter preference bonus and utilize the entrance mask. - core: adjust optimizer bounds and scoring weights to favor larger diameter screws and improve convergence. - xfr_cbt_native: implement `render_xray_projections` to generate synthetic AP and lateral X-ray images for visualization. - visualization: enhance `render_bone_figure` with semi-transparent spinous process rendering and improved depth sorting for screws. - xfr_debug: improve level detection to support arbitrary lumbar levels (L1-L9) and add safe volume-level cleanup for CBT writing. - config: update allowed diameters and lengths constants.
242 lines
No EOL
9.1 KiB
Python
242 lines
No EOL
9.1 KiB
Python
#!/home/xfr/.conda/envs/cbt/bin/python
|
||
"""
|
||
從 optimizer 存的 output.csv 重新渲染既有螺絲模式四視角圖(X-ray 圖)。
|
||
|
||
用途:render_bone_figure 改了純顯示層(例如棘突改半透明 + 依深度正確遮擋
|
||
螺絲路徑)之後,不需要重跑優化,直接從 CSV 的 best_position 重新出圖。
|
||
原圖先備份到 Output/{date}/backup_opaque_spinous/{vol}/,再重渲染同檔名覆蓋。
|
||
|
||
mask 取法與 xfr_debug / xfr_plot_level 相同(level_file_path:
|
||
binary_sdf 優先、缺則 binary;cortical 同),spacing 讀自 mask 檔。
|
||
|
||
Usage:
|
||
python xfr_rerender_spinous.py # 預設 20260912,多 GPU 並行
|
||
python xfr_rerender_spinous.py 20260911 # 指定日期
|
||
python xfr_rerender_spinous.py 20260912 0001 # 只該 volume(全名或末段)
|
||
python xfr_rerender_spinous.py --dry-run 20260912 # 只列出會重新渲染的圖
|
||
python xfr_rerender_spinous.py --cpus 20260912 # 純 CPU 循序
|
||
"""
|
||
|
||
import argparse
|
||
import csv
|
||
import multiprocessing
|
||
import os
|
||
import re
|
||
import shutil
|
||
import sys
|
||
|
||
import numpy as np
|
||
import SimpleITK as sitk
|
||
|
||
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
|
||
|
||
from imaging.transforms import level_file_path
|
||
from visualization.res_bone_figure import render_bone_figure
|
||
|
||
OUTPUT_BASE = '/mnt/1248/open2/cyrou/Output'
|
||
MASK_DIR = '/mnt/1248/open2/cyrou/CBT/Seg/Resample/standardized-xfr'
|
||
BACKUP = 'backup_opaque_spinous'
|
||
|
||
# {level}_{way}_L{d}_{l}_R{d}_{l}_{swarm}_{iter}.png(單側跑法該側可為空)
|
||
PNG_RE = re.compile(
|
||
r'^(?P<level>[A-Z]\d+?)_(?P<way>CBT|TPS)'
|
||
r'_L(?P<dl>\d+(?:\.\d+)?|)_(?P<ll>\d+(?:\.\d+)?|)'
|
||
r'_R(?P<dr>\d+(?:\.\d+)?|)_(?P<lr>\d+(?:\.\d+)?|)'
|
||
r'_(?P<sw>\d+|)_(?P<it>\d+|)\.png$')
|
||
|
||
|
||
def _f(s):
|
||
s = (s or '').strip()
|
||
return float(s) if s else None
|
||
|
||
|
||
def _parse_pos(s):
|
||
"""CSV '(x, y, z)' -> (z, y, x)(best_position 慣例)。"""
|
||
v = re.findall(r'[-+]?\d*\.?\d+(?:[eE][-+]?\d+)?', s or '')
|
||
if len(v) < 3:
|
||
return None
|
||
x, y, z = (float(t) for t in v[:3])
|
||
return z, y, x
|
||
|
||
|
||
def _pos(row):
|
||
if row is None:
|
||
return None
|
||
p = _parse_pos(row.get('Position_XYZ'))
|
||
if p is None:
|
||
return None
|
||
az, alt = _f(row.get('Raw_Azimuth')), _f(row.get('Raw_Altitude'))
|
||
if az is None or alt is None:
|
||
return None
|
||
return (p[0], p[1], p[2], az, alt)
|
||
|
||
|
||
def read_csv_rows(vol_out_dir):
|
||
path = os.path.join(vol_out_dir, 'output.csv')
|
||
if not os.path.isfile(path):
|
||
return None
|
||
with open(path, newline='') as f:
|
||
return [r for r in csv.DictReader(f) if any(c.strip() for c in r.values())]
|
||
|
||
|
||
def find_row(rows, side, d, l):
|
||
if d is None or l is None:
|
||
return None
|
||
for r in rows:
|
||
if r.get('Side') != side:
|
||
continue
|
||
rd, rl = _f(r.get('Diameter')), _f(r.get('Length'))
|
||
if rd is not None and rl is not None \
|
||
and abs(rd - d) < 1e-6 and abs(rl - l) < 1e-6:
|
||
return r
|
||
return None
|
||
|
||
|
||
def collect_figs(date_dir, vol_filter=None):
|
||
"""回傳 (tasks, skipped):tasks 為可重新渲染的圖(dict),skipped 為 (原因, 路徑)。"""
|
||
vols = sorted(d for d in os.listdir(date_dir)
|
||
if os.path.isdir(os.path.join(date_dir, d)) and d != BACKUP)
|
||
if vol_filter:
|
||
vols = [v for v in vols
|
||
if v == vol_filter or v.rsplit('.', 1)[-1] == vol_filter]
|
||
if not vols:
|
||
raise SystemExit(f'volume not found: {vol_filter}')
|
||
tasks, skipped = [], []
|
||
for vol in vols:
|
||
vdir = os.path.join(date_dir, vol)
|
||
rows = read_csv_rows(vdir) or []
|
||
for name in sorted(os.listdir(vdir)):
|
||
m = PNG_RE.match(name)
|
||
if not m:
|
||
continue
|
||
png = os.path.join(vdir, name)
|
||
level, way = m['level'], m['way']
|
||
d_l, l_l = _f(m['dl']), _f(m['ll'])
|
||
d_r, l_r = _f(m['dr']), _f(m['lr'])
|
||
row_l = find_row(rows, 'L', d_l, l_l)
|
||
row_r = find_row(rows, 'R', d_r, l_r)
|
||
if (d_l is not None and row_l is None) or (d_r is not None and row_r is None) \
|
||
or (d_l is None and d_r is None):
|
||
skipped.append((vol, name, 'CSV 無對應 L/R 行'))
|
||
continue
|
||
mask_dir = os.path.join(MASK_DIR, vol)
|
||
sdf = level_file_path(mask_dir, level, 'binary_sdf')
|
||
binary_path = sdf if os.path.exists(sdf) \
|
||
else level_file_path(mask_dir, level, 'binary')
|
||
if not os.path.exists(binary_path):
|
||
skipped.append((vol, name, f'無 bone mask: {binary_path}'))
|
||
continue
|
||
cortical_path = level_file_path(mask_dir, level, 'cortical')
|
||
row_any = row_l or row_r
|
||
tt = _f(row_any.get('Total_Time'))
|
||
tt = None if (tt is None or not np.isfinite(tt)) else tt
|
||
tasks.append({
|
||
'vol': vol, 'name': name, 'png': png, 'level': level, 'way': way,
|
||
'd_l': d_l, 'l_l': l_l, 'd_r': d_r, 'l_r': l_r,
|
||
'pos_l': _pos(row_l), 'pos_r': _pos(row_r),
|
||
'binary': binary_path,
|
||
'cortical': cortical_path if os.path.exists(cortical_path) else None,
|
||
'swarm': int(_f(m['sw']) or 0), 'iter': int(_f(m['it']) or 0),
|
||
'time': tt,
|
||
})
|
||
return tasks, skipped
|
||
|
||
|
||
def _run_batch(batch):
|
||
"""一個 worker 固定綁一個 GPU,依序渲染分到的圖(CUDA_VISIBLE_DEVICES
|
||
須在 torch 首次 init CUDA 前設好,故綁定後不再變)。"""
|
||
gpu, g_tasks = batch
|
||
os.environ['CUDA_VISIBLE_DEVICES'] = str(gpu)
|
||
os.environ.setdefault('OMP_NUM_THREADS', '4')
|
||
import matplotlib
|
||
matplotlib.use('Agg')
|
||
return [render_one(t) for t in g_tasks]
|
||
|
||
|
||
def render_one(task):
|
||
import torch
|
||
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
|
||
img = sitk.ReadImage(task['binary'], sitk.sitkUInt8)
|
||
spacing = list(img.GetSpacing())
|
||
|
||
# 備份原圖(同目錄同檔名會重複跑時,備份檔保留第一份)
|
||
backup_dir = os.path.join(os.path.dirname(os.path.dirname(task['png'])), BACKUP, task['vol'])
|
||
bpath = os.path.join(backup_dir, task['name'])
|
||
os.makedirs(backup_dir, exist_ok=True)
|
||
if not os.path.exists(bpath):
|
||
shutil.move(task['png'], bpath)
|
||
|
||
try:
|
||
out = render_bone_figure(
|
||
task['vol'], task['level'], task['binary'], task['cortical'],
|
||
base_folder=OUTPUT_BASE, spacing=spacing, way=task['way'],
|
||
best_position_l=task['pos_l'], best_position_r=task['pos_r'],
|
||
diameter_l=task['d_l'], length_l=task['l_l'],
|
||
diameter_r=task['d_r'], length_r=task['l_r'],
|
||
image2_path=None, device=device,
|
||
swarm_size=task['swarm'], max_iter=task['iter'], total_time=task['time'],
|
||
write_csv=False, output_path=task['png'])
|
||
except Exception as e:
|
||
if not os.path.exists(task['png']) and os.path.exists(bpath):
|
||
shutil.move(bpath, task['png'])
|
||
return (task['vol'], task['name'], 'fail', f'{type(e).__name__}: {e}')
|
||
if out is None:
|
||
if os.path.exists(bpath):
|
||
shutil.move(bpath, task['png'])
|
||
return (task['vol'], task['name'], 'fail', 'render 回傳 None(無/空遮罩)')
|
||
if out != task['png']:
|
||
os.replace(out, task['png'])
|
||
return (task['vol'], task['name'], 'ok', f'{out} ({device})')
|
||
|
||
|
||
def main():
|
||
ap = argparse.ArgumentParser(description='從 output.csv 重新渲染螺絲模式四視角圖')
|
||
ap.add_argument('date', nargs='?', default='20260912')
|
||
ap.add_argument('volume', nargs='?', default=None)
|
||
ap.add_argument('--dry-run', action='store_true')
|
||
ap.add_argument('--cpus', action='store_true', help='CPU 循序(不佔 GPU)')
|
||
args = ap.parse_args()
|
||
|
||
date_dir = os.path.join(OUTPUT_BASE, args.date)
|
||
if not os.path.isdir(date_dir):
|
||
raise SystemExit(f'no such date dir: {date_dir}')
|
||
tasks, skipped = collect_figs(date_dir, args.volume)
|
||
print(f'{args.date}: {len(tasks)} figure(s) to re-render, {len(skipped)} skip(s)')
|
||
for vol, name, why in skipped:
|
||
print(f' [skip] {vol}/{name}: {why}')
|
||
if args.dry_run:
|
||
for t in tasks:
|
||
print(f" [dry] {t['vol']}/{t['name']} "
|
||
f"L=({t['d_l']},{t['l_l']}) R=({t['d_r']},{t['l_r']})")
|
||
return
|
||
|
||
if not tasks:
|
||
return
|
||
|
||
if args.cpus or not _cuda_count():
|
||
os.environ['CUDA_VISIBLE_DEVICES'] = ''
|
||
for i, t in enumerate(tasks, 1):
|
||
r = render_one(t)
|
||
print(f'[{i:3d}/{len(tasks)}] {r[2]:4s} {r[0]}/{r[1]} {r[3]}', flush=True)
|
||
else:
|
||
n = min(_cuda_count(), 4)
|
||
ctx = multiprocessing.get_context('spawn')
|
||
batches = [(i, tasks[i::n]) for i in range(n)]
|
||
with ctx.Pool(n) as pool:
|
||
for results in pool.map(_run_batch, batches):
|
||
for vol, name, status, detail in results:
|
||
print(f'[{status:4s}] {vol}/{name} {detail}', flush=True)
|
||
print('=' * 60)
|
||
print('Done.')
|
||
|
||
|
||
def _cuda_count():
|
||
try:
|
||
import torch
|
||
return torch.cuda.device_count()
|
||
except Exception:
|
||
return 0
|
||
|
||
|
||
if __name__ == '__main__':
|
||
main() |