CoolFace
Modelpublic

BAAI/SegVol

sourceHugging Faceupdated 2y agoView on Hugging Face
13likes4.7kdownloads
README.md192 linesDownload Raw Back to root
1 2![image/jpeg](https://cdn-uploads.huggingface.co/production/uploads/6565b54a9bf6665f10f75441/no60wyvKDTD-WV3pCt2P5.jpeg)3 4Language: [EN / ZH]5 6The SegVol is a universal and interactive model for volumetric medical image segmentation. SegVol accepts point, box, and text prompts while output volumetric segmentation. By training on 90k unlabeled Computed Tomography (CT) volumes and 6k labeled CTs, this foundation model supports the segmentation of over 200 anatomical categories.7 8SegVol是用于体积医学图像分割的通用交互式模型,可以使用点,框和文本作为prompt驱动模型,输出分割结果。9 10通过在90k个无标签CT和6k个有标签CT上进行训练,该基础模型支持对200多个解剖类别进行分割。11 12[**Paper**](https://arxiv.org/abs/2311.13385), [**Code**](https://github.com/BAAI-DCAI/SegVol) 和 [**Demo**](https://huggingface.co/spaces/BAAI/SegVol) 已发布。13 14**Keywords**: 3D medical SAM, volumetric image segmentation15 16## Quicktart17 18### Requirements19```bash20conda create -n segvol_transformers python=3.821conda activate segvol_transformers22```23[pytorch v1.11.0](https://pytorch.org/get-started/previous-versions/) or higher version is required. Please also install the following support packages:24 25需要 [pytorch v1.11.0](https://pytorch.org/get-started/previous-versions/) 或更高版本。另外请安装如下支持包:26 27```bash28pip install 'monai[all]==0.9.0'29pip install einops==0.6.130pip install transformers==4.18.031pip install matplotlib32```33 34### Test script35 36```python37from transformers import AutoModel, AutoTokenizer38import torch39import os40 41# get device42device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")43 44# load model45clip_tokenizer = AutoTokenizer.from_pretrained("BAAI/SegVol")46model = AutoModel.from_pretrained("BAAI/SegVol", trust_remote_code=True, test_mode=True)47model.model.text_encoder.tokenizer = clip_tokenizer48model.eval()49model.to(device)50print('model load done')51 52# set case path53ct_path = 'path/to/Case_image_00001_0000.nii.gz'54gt_path = 'path/to/Case_label_00001.nii.gz'55 56# set categories, corresponding to the unique values(1, 2, 3, 4, ...) in ground truth mask57categories = ["liver", "kidney", "spleen", "pancreas"]58 59# generate npy data format60ct_npy, gt_npy = model.processor.preprocess_ct_gt(ct_path, gt_path, category=categories)61# IF you have download our 25 processed datasets, you can skip to here with the processed ct_npy, gt_npy files62 63# go through zoom_transform to generate zoomout & zoomin views64data_item = model.processor.zoom_transform(ct_npy, gt_npy)65 66# add batch dim manually67data_item['image'], data_item['label'], data_item['zoom_out_image'], data_item['zoom_out_label'] = \68data_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)69 70# take liver as the example71cls_idx = 072 73# text prompt74text_prompt = [categories[cls_idx]]75 76# point prompt77point_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 dim78 79# bbox prompt80bbox_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 dim81 82print('prompt done')83 84# segvol test forward85# use_zoom: use zoom-out-zoom-in86# point_prompt_group: use point prompt87# bbox_prompt_group: use bbox prompt88# text_prompt: use text prompt89logits_mask = model.forward_test(image=data_item['image'],90      zoomed_image=data_item['zoom_out_image'],91      # point_prompt_group=[point_prompt, point_prompt_map],92      bbox_prompt_group=[bbox_prompt, bbox_prompt_map],93      text_prompt=text_prompt,94      use_zoom=True95      )96 97# cal dice score98dice = model.processor.dice_score(logits_mask[0][0], data_item['label'][0][cls_idx], device)99print(dice)100 101# save prediction as nii.gz file102save_path='./Case_preds_00001.nii.gz'103model.processor.save_preds(ct_path, save_path, logits_mask[0][0], 104                           start_coord=data_item['foreground_start_coord'], 105                           end_coord=data_item['foreground_end_coord'])106print('done')107```108 109### Training script110 111```python112from transformers import AutoModel, AutoTokenizer113import torch114import os115 116# get device117device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")118 119# load model120clip_tokenizer = AutoTokenizer.from_pretrained("BAAI/SegVol")121model = AutoModel.from_pretrained("BAAI/SegVol", trust_remote_code=True, test_mode=False)122model.model.text_encoder.tokenizer = clip_tokenizer123model.train()124model.to(device)125print('model load done')126 127# set case path128ct_path = 'path/to/Case_image_00001_0000.nii.gz'129gt_path = 'path/to/Case_label_00001.nii.gz'130 131# set categories, corresponding to the unique values(1, 2, 3, 4, ...) in ground truth mask132categories = ["liver", "kidney", "spleen", "pancreas"]133 134# generate npy data format135ct_npy, gt_npy = model.processor.preprocess_ct_gt(ct_path, gt_path, category=categories)136# IF you have download our 25 processed datasets, you can skip to here with the processed ct_npy, gt_npy files137 138# go through train transform139data_item = model.processor.train_transform(ct_npy, gt_npy)140 141# training example142# add batch dim manually143image, gt3D = data_item["image"].unsqueeze(0).to(device), data_item["label"].unsqueeze(0).to(device) # add batch dim144 145loss_step_avg = 0146for cls_idx in range(len(categories)):147    # optimizer.zero_grad()148    organs_cls = categories[cls_idx]149    labels_cls = gt3D[:, cls_idx]150    loss = model.forward_train(image, train_organs=organs_cls, train_labels=labels_cls)151    loss_step_avg += loss.item()152    loss.backward()153    # optimizer.step()154 155loss_step_avg /= len(categories)156print(f'AVG loss {loss_step_avg}')157 158# save ckpt159model.save_pretrained('./ckpt')160```161 162### Start with M3D-Seg dataset163 164We have released 25 open source datasets(M3D-Seg) for training SegVol, and these preprocessed data have been uploaded to [ModelScope](https://www.modelscope.cn/datasets/GoodBaiBai88/M3D-Seg/summary) and [HuggingFace](https://huggingface.co/datasets/GoodBaiBai88/M3D-Seg). 165You can use the following script to easily load cases and insert them into Test script and Training script.166 167我们已经发布了用于训练SegVol的25个开源数据集(M3D-Seg),并将预处理后的数据上传到了[ModelScope](https://www.modelscope.cn/datasets/GoodBaiBai88/M3D-Seg/summary)和[HuggingFace](https://huggingface.co/datasets/GoodBaiBai88/M3D-Seg)。 168您可以使用下面的script方便地载入,并插入到Test script和Training script中。169 170```python171import json, os172M3D_Seg_path = 'path/to/M3D-Seg'173 174# select a dataset175dataset_code = '0000'176 177# load json dict178json_path = os.path.join(M3D_Seg_path, dataset_code, dataset_code + '.json')179with open(json_path, 'r') as f:180    dataset_dict = json.load(f)181 182# get a case183ct_path = os.path.join(M3D_Seg_path, dataset_dict['train'][0]['image'])184gt_path = os.path.join(M3D_Seg_path, dataset_dict['train'][0]['label'])185 186# get categories187categories_dict = dataset_dict['labels']188categories = [x for _, x in categories_dict.items() if x != "background"]189 190# load npy data format191ct_npy, gt_npy = model.processor.load_uniseg_case(ct_path, gt_path)192```