"""Chạy lại pipeline (Model A → crop → Model B) trên 1 ảnh và viz. Output ra /tmp.""" import argparse, json from pathlib import Path import numpy as np import matplotlib; matplotlib.use('Agg') import matplotlib.pyplot as plt import matplotlib.patches as mpatches from PIL import Image import torch import sys ROOT = Path(__file__).resolve().parents[2] # ai_drawing/ sys.path.insert(0, str(ROOT)) sys.path.insert(0, str(ROOT / 'shared')) from evaluate_view_pipeline import ( load_model, run_view_pipeline, run_baseline, decode_rle, CKPT_VIEW_A, CKPT_VIEW_B, ID2LABEL_VIEW, _id2label_for_ckpt, SE_VIEW_A, LE_VIEW_A, SE_VIEW_B, LE_VIEW_B, SCORE_THRESH_A, SCORE_THRESH_B, BASE_DIR, _PALETTE, ) GT_JSON = BASE_DIR / 'data' / 'whole_images_coco_view' / 'test.json' IMG_DIR = BASE_DIR / 'data' / 'whole_images' / 'test' def main(): ap = argparse.ArgumentParser() ap.add_argument('--match', default='R013') ap.add_argument('--out', default='/tmp/REPRED_R013.png') args = ap.parse_args() gt = json.load(open(GT_JSON)) im = next(i for i in gt['images'] if args.match in i['file_name']) fn = im['file_name']; H, W = im['height'], im['width'] print(f"Image: {fn} {W}x{H}") device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') model_a, proc_a = load_model(CKPT_VIEW_A, SE_VIEW_A, LE_VIEW_A, device, ID2LABEL_VIEW) id2b = _id2label_for_ckpt(CKPT_VIEW_B) model_b, proc_b = load_model(CKPT_VIEW_B, SE_VIEW_B, LE_VIEW_B, device, id2b) print(f"Model B id2label = {id2b}") pil = Image.open(IMG_DIR / fn).convert('RGB') dets_a = run_baseline(pil, model_a, proc_a, device, (H, W), SCORE_THRESH_A) dets_pipe, timing = run_view_pipeline(pil, model_a, proc_a, model_b, proc_b, device, (H, W), SCORE_THRESH_A, SCORE_THRESH_B) n_gt = len([a for a in gt['annotations'] if a['image_id'] == im['id']]) print(f"Model A dets={len(dets_a)} Pipeline dets={len(dets_pipe)} GT={n_gt} " f"t_total={timing['t_total']:.1f}s n_crops={timing['n_crops']}") img = np.array(pil) DPI = 100; s = W / 1000.0 fig = plt.figure(figsize=(W / DPI, H / DPI), dpi=DPI) ax = fig.add_axes([0, 0, 1, 1]); ax.imshow(img); ax.axis('off') ax.set_xlim(0, W); ax.set_ylim(H, 0) ov = np.zeros((H, W, 4), np.float32) for idx, d in enumerate(dets_pipe): m = decode_rle(d['segmentation']).astype(bool) r, g, b = _PALETTE[idx % len(_PALETTE)] ov[m] = [r, g, b, 0.45] ax.imshow(ov) for d in dets_a: x, y, w, h = d['bbox'] ax.add_patch(plt.Rectangle((x, y), w, h, lw=2 * s, ls='--', edgecolor='yellow', facecolor='none')) for d in dets_pipe: x, y, w, h = d['bbox'] ax.add_patch(plt.Rectangle((x, y), w, h, lw=2 * s, edgecolor='deepskyblue', facecolor='none')) for a in [a for a in gt['annotations'] if a['image_id'] == im['id']]: x, y, w, h = a['bbox'] ax.add_patch(plt.Rectangle((x, y), w, h, lw=2 * s, edgecolor='lime', facecolor='none')) ax.legend(handles=[ mpatches.Patch(facecolor='none', edgecolor='yellow', ls='--', label=f'Model A box ({len(dets_a)})'), mpatches.Patch(facecolor='none', edgecolor='deepskyblue', label=f'Pipeline/Model B box ({len(dets_pipe)})'), mpatches.Patch(facecolor='steelblue', alpha=0.45, label='Model B mask'), mpatches.Patch(facecolor='none', edgecolor='lime', label=f'GT box ({n_gt})'), ], loc='upper right', fontsize=9 * s) fig.savefig(args.out, dpi=DPI); plt.close(fig) print(f"Saved → {args.out} size={Image.open(args.out).size}") if __name__ == '__main__': main()