Views
No views yet

1conda create -n segvol_transformers python=3.8
2conda activate segvol_transformers1pip install 'monai[all]==0.9.0'
2pip install einops==0.6.1
3pip install transformers==4.18.0
4pip install matplotlib1from transformers import AutoModel, AutoTokenizer
2import torch
3import os
4
5# get device
6device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
7
8# load model
9clip_tokenizer = AutoTokenizer.from_pretrained("BAAI/SegVol")
10model = AutoModel.from_pretrained("BAAI/SegVol", trust_remote_code=True, test_mode=True)
11model.model.text_encoder.tokenizer = clip_tokenizer
12model.eval()
13model.to(device)
14print('model load done')
15
16# set case path
17ct_path = 'path/to/Case_image_00001_0000.nii.gz'
18gt_path = 'path/to/Case_label_00001.nii.gz'
19
20# set categories, corresponding to the unique values(1, 2, 3, 4, ...) in ground truth mask
21categories = ["liver", "kidney", "spleen", "pancreas"]
22
23# generate npy data format
24ct_npy, gt_npy = model.processor.preprocess_ct_gt(ct_path, gt_path, category=categories)
25# IF you have download our 25 processed datasets, you can skip to here with the processed ct_npy, gt_npy files
26
27# go through zoom_transform to generate zoomout & zoomin views
28data_item = model.processor.zoom_transform(ct_npy, gt_npy)
29
30# add batch dim manually
31data_item['image'], data_item['label'], data_item['zoom_out_image'], data_item['zoom_out_label'] = \
32data_item['image'].unsqueeze(0).to(device), data_item['label'].unsqueeze(0).to(device), data_item['zoom_out_image'].unsqueeze(0).to(device), data_item['zoom_out_label'].unsqueeze(0).to(device)
33
34# take liver as the example
35cls_idx = 0
36
37# text prompt
38text_prompt = [categories[cls_idx]]
39
40# point prompt
41point_prompt, point_prompt_map = model.processor.point_prompt_b(data_item['zoom_out_label'][0][cls_idx], device=device) # inputs w/o batch dim, outputs w batch dim
42
43# bbox prompt
44bbox_prompt, bbox_prompt_map = model.processor.bbox_prompt_b(data_item['zoom_out_label'][0][cls_idx], device=device) # inputs w/o batch dim, outputs w batch dim
45
46print('prompt done')
47
48# segvol test forward
49# use_zoom: use zoom-out-zoom-in
50# point_prompt_group: use point prompt
51# bbox_prompt_group: use bbox prompt
52# text_prompt: use text prompt
53logits_mask = model.forward_test(image=data_item['image'],
54 zoomed_image=data_item['zoom_out_image'],
55 # point_prompt_group=[point_prompt, point_prompt_map],
56 bbox_prompt_group=[bbox_prompt, bbox_prompt_map],
57 text_prompt=text_prompt,
58 use_zoom=True
59 )
60
61# cal dice score
62dice = model.processor.dice_score(logits_mask[0][0], data_item['label'][0][cls_idx], device)
63print(dice)
64
65# save prediction as nii.gz file
66save_path='./Case_preds_00001.nii.gz'
67model.processor.save_preds(ct_path, save_path, logits_mask[0][0],
68 start_coord=data_item['foreground_start_coord'],
69 end_coord=data_item['foreground_end_coord'])
70print('done')1from transformers import AutoModel, AutoTokenizer
2import torch
3import os
4
5# get device
6device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
7
8# load model
9clip_tokenizer = AutoTokenizer.from_pretrained("BAAI/SegVol")
10model = AutoModel.from_pretrained("BAAI/SegVol", trust_remote_code=True, test_mode=False)
11model.model.text_encoder.tokenizer = clip_tokenizer
12model.train()
13model.to(device)
14print('model load done')
15
16# set case path
17ct_path = 'path/to/Case_image_00001_0000.nii.gz'
18gt_path = 'path/to/Case_label_00001.nii.gz'
19
20# set categories, corresponding to the unique values(1, 2, 3, 4, ...) in ground truth mask
21categories = ["liver", "kidney", "spleen", "pancreas"]
22
23# generate npy data format
24ct_npy, gt_npy = model.processor.preprocess_ct_gt(ct_path, gt_path, category=categories)
25# IF you have download our 25 processed datasets, you can skip to here with the processed ct_npy, gt_npy files
26
27# go through train transform
28data_item = model.processor.train_transform(ct_npy, gt_npy)
29
30# training example
31# add batch dim manually
32image, gt3D = data_item["image"].unsqueeze(0).to(device), data_item["label"].unsqueeze(0).to(device) # add batch dim
33
34loss_step_avg = 0
35for cls_idx in range(len(categories)):
36 # optimizer.zero_grad()
37 organs_cls = categories[cls_idx]
38 labels_cls = gt3D[:, cls_idx]
39 loss = model.forward_train(image, train_organs=organs_cls, train_labels=labels_cls)
40 loss_step_avg += loss.item()
41 loss.backward()
42 # optimizer.step()
43
44loss_step_avg /= len(categories)
45print(f'AVG loss {loss_step_avg}')
46
47# save ckpt
48model.save_pretrained('./ckpt')1import json, os
2M3D_Seg_path = 'path/to/M3D-Seg'
3
4# select a dataset
5dataset_code = '0000'
6
7# load json dict
8json_path = os.path.join(M3D_Seg_path, dataset_code, dataset_code + '.json')
9with open(json_path, 'r') as f:
10 dataset_dict = json.load(f)
11
12# get a case
13ct_path = os.path.join(M3D_Seg_path, dataset_dict['train'][0]['image'])
14gt_path = os.path.join(M3D_Seg_path, dataset_dict['train'][0]['label'])
15
16# get categories
17categories_dict = dataset_dict['labels']
18categories = [x for _, x in categories_dict.items() if x != "background"]
19
20# load npy data format
21ct_npy, gt_npy = model.processor.load_uniseg_case(ct_path, gt_path)