-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathSAM_cam.py
More file actions
84 lines (60 loc) · 2.66 KB
/
Copy pathSAM_cam.py
File metadata and controls
84 lines (60 loc) · 2.66 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
import numpy as np
from MedSAM.segment_anything import sam_model_registry
from MedSAM.demo import BboxPromptDemo
import os
from CAM.utils import GradCAM
from CAM.vit_model import vit_base_patch16_224_in21k as create_model
import argparse
import CAM.main_vit
from CAM.main_vit import show_mask,dice_coeff
import logging
import time
from CAM.utils import scoremap2bbox
from Utils import *
def main(opt,logger):
device = "cuda:0"
#CAM模型导入
MedSAM_CKPT_PATH = "C:/Users/Administrator/Desktop/MedSAM-main/work_dir/MedSAM/medsam_vit_b.pth"
medsam_model = sam_model_registry['vit_b'](checkpoint=MedSAM_CKPT_PATH)
medsam_model = medsam_model.to(device)
medsam_model.eval()
bbox_prompt_demo = BboxPromptDemo(medsam_model)
for i,img_name in enumerate(os.listdir(opt.img_path),1):
img_data = np.load(os.path.join(opt.img_path, img_name))["arr_0"]
label_data = np.load(os.path.join(opt.label_path, img_name))["arr_0"]
label_data[label_data > 0] = 1
file_cam=np.load(os.path.join(opt.cam_path,img_name))["arr_0"]
grad_cam=file_cam['original_cam']
caa_grad_cam=file_cam['caa_cam']
def loadLogger(args):
logger = logging.getLogger()
logger.setLevel(logging.INFO)
formatter = logging.Formatter(fmt="[ %(asctime)s ] %(message)s",
datefmt="%a %b %d %H:%M:%S %Y")
sHandler = logging.StreamHandler()
sHandler.setFormatter(formatter)
logger.addHandler(sHandler)
if not args.not_save:
work_dir = os.path.join(args.work_dir,
time.strftime("%Y.%m.%dT%H %M %S", time.localtime()))
if not os.path.exists(work_dir):
os.makedirs(work_dir)
fHandler = logging.FileHandler(work_dir + '/log.txt', mode='w')
fHandler.setLevel(logging.DEBUG)
fHandler.setFormatter(formatter)
logger.addHandler(fHandler)
return logger
if __name__ == '__main__':
parser = argparse.ArgumentParser()
parser.add_argument('--device', default='cuda:0', help='device id (i.e. 0 or 0,1 or cpu)')
parser.add_argument('--num_classes', type=int, default=2)
parser.add_argument('--img_path',default="F:/brats/val/yes")
parser.add_argument('--label_path',default="F:/brats/val/label")
parser.add_argument('--not-save', default=False, action='store_true',
help='If yes, only output log to terminal.')
parser.add_argument('--cam_path',default="F:/brats/val/cam")
parser.add_argument('--work-dir', default='./work_dir',
help='the work folder for storing results')
opt = parser.parse_args()
logger = loadLogger(opt)
main(opt,logger)