Foundation model for image segmentation with zero-shot transfer. Use when you need to segment any object in images using points, boxes, or masks as prompts, or automatically generate all object masks in an image.
1.0.0
Orchestra Research
MIT
segment-anything
transformers>=4.30.0
torch>=1.7.0
hermes
tags
Multimodal
Image Segmentation
Computer Vision
SAM
Zero-Shot
Segment Anything Model (SAM)
Comprehensive guide to using Meta AI's Segment Anything Model for zero-shot image segmentation.
When to use SAM
Use SAM when:
Need to segment any object in images without task-specific training
Building interactive annotation tools with point/box prompts
Generating training data for other vision models
Need zero-shot transfer to new image domains
Building object detection/segmentation pipelines
Processing medical, satellite, or domain-specific images
Key features:
Zero-shot segmentation: Works on any image domain without fine-tuning
Flexible prompts: Points, bounding boxes, or previous masks
Automatic segmentation: Generate all object masks automatically
High quality: Trained on 1.1 billion masks from 11 million images
Multiple model sizes: ViT-B (fastest), ViT-L, ViT-H (most accurate)
ONNX export: Deploy in browsers and edge devices
Use alternatives instead:
YOLO/Detectron2: For real-time object detection with classes
Mask2Former: For semantic/panoptic segmentation with categories
GroundingDINO + SAM: For text-prompted segmentation
SAM 2: For video segmentation tasks
Quick start
Installation
# From GitHub
pip install git+https://github.com/facebookresearch/segment-anything.git
# Optional dependencies
pip install opencv-python pycocotools matplotlib
# Or use HuggingFace transformers
pip install transformers
importnumpyasnpfromsegment_anythingimportsam_model_registry,SamPredictor# Load modelsam=sam_model_registry["vit_h"](checkpoint="sam_vit_h_4b8939.pth")sam.to(device="cuda")# Create predictorpredictor=SamPredictor(sam)# Set image (computes embeddings once)image=cv2.imread("image.jpg")image=cv2.cvtColor(image,cv2.COLOR_BGR2RGB)predictor.set_image(image)# Predict with point promptsinput_point=np.array([[500,375]])# (x, y) coordinatesinput_label=np.array([1])# 1 = foreground, 0 = backgroundmasks,scores,logits=predictor.predict(point_coords=input_point,point_labels=input_label,multimask_output=True# Returns 3 mask options)# Select best maskbest_mask=masks[np.argmax(scores)]
HuggingFace Transformers
importtorchfromPILimportImagefromtransformersimportSamModel,SamProcessor# Load model and processormodel=SamModel.from_pretrained("facebook/sam-vit-huge")processor=SamProcessor.from_pretrained("facebook/sam-vit-huge")model.to("cuda")# Process image with point promptimage=Image.open("image.jpg")input_points=[[[450,600]]]# Batch of pointsinputs=processor(image,input_points=input_points,return_tensors="pt")inputs={k:v.to("cuda")fork,vininputs.items()}# Generate maskswithtorch.no_grad():outputs=model(**inputs)# Post-process masks to original sizemasks=processor.image_processor.post_process_masks(outputs.pred_masks.cpu(),inputs["original_sizes"].cpu(),inputs["reshaped_input_sizes"].cpu())
# Single foreground pointinput_point=np.array([[500,375]])input_label=np.array([1])masks,scores,logits=predictor.predict(point_coords=input_point,point_labels=input_label,multimask_output=True)# Multiple points (foreground + background)input_points=np.array([[500,375],[600,400],[450,300]])input_labels=np.array([1,1,0])# 2 foreground, 1 backgroundmasks,scores,logits=predictor.predict(point_coords=input_points,point_labels=input_labels,multimask_output=False# Single mask when prompts are clear)
# Box + points for precise controlmasks,scores,logits=predictor.predict(point_coords=np.array([[500,375]]),point_labels=np.array([1]),box=np.array([400,300,700,600]),multimask_output=False)
Iterative refinement
# Initial predictionmasks,scores,logits=predictor.predict(point_coords=np.array([[500,375]]),point_labels=np.array([1]),multimask_output=True)# Refine with additional point using previous maskmasks,scores,logits=predictor.predict(point_coords=np.array([[500,375],[550,400]]),point_labels=np.array([1,0]),# Add background pointmask_input=logits[np.argmax(scores)][None,:,:],# Use best maskmultimask_output=False)
Automatic mask generation
Basic automatic segmentation
fromsegment_anythingimportSamAutomaticMaskGenerator# Create generatormask_generator=SamAutomaticMaskGenerator(sam)# Generate all masksmasks=mask_generator.generate(image)# Each mask contains:# - segmentation: binary mask# - bbox: [x, y, w, h]# - area: pixel count# - predicted_iou: quality score# - stability_score: robustness score# - point_coords: generating point
Customized generation
mask_generator=SamAutomaticMaskGenerator(model=sam,points_per_side=32,# Grid density (more = more masks)pred_iou_thresh=0.88,# Quality thresholdstability_score_thresh=0.95,# Stability thresholdcrop_n_layers=1,# Multi-scale cropscrop_n_points_downscale_factor=2,min_mask_region_area=100,# Remove tiny masks)masks=mask_generator.generate(image)
Filtering masks
# Sort by area (largest first)masks=sorted(masks,key=lambdax:x['area'],reverse=True)# Filter by predicted IoUhigh_quality=[mforminmasksifm['predicted_iou']>0.9]# Filter by stability scorestable_masks=[mforminmasksifm['stability_score']>0.95]
Batched inference
Multiple images
# Process multiple images efficientlyimages=[cv2.imread(f"image_{i}.jpg")foriinrange(10)]all_masks=[]forimageinimages:predictor.set_image(image)masks,_,_=predictor.predict(point_coords=np.array([[500,375]]),point_labels=np.array([1]),multimask_output=True)all_masks.append(masks)
Multiple prompts per image
# Process multiple prompts efficiently (one image encoding)predictor.set_image(image)# Batch of point promptspoints=[np.array([[100,100]]),np.array([[200,200]]),np.array([[300,300]])]all_masks=[]forpointinpoints:masks,scores,_=predictor.predict(point_coords=point,point_labels=np.array([1]),multimask_output=True)all_masks.append(masks[np.argmax(scores)])
importonnxruntime# Load ONNX modelort_session=onnxruntime.InferenceSession("sam_onnx.onnx")# Run inference (image embeddings computed separately)masks=ort_session.run(None,{"image_embeddings":image_embeddings,"point_coords":point_coords,"point_labels":point_labels,"mask_input":np.zeros((1,1,256,256),dtype=np.float32),"has_mask_input":np.array([0],dtype=np.float32),"orig_im_size":np.array([h,w],dtype=np.float32)})
Common workflows
Workflow 1: Annotation tool
importcv2# Load modelpredictor=SamPredictor(sam)predictor.set_image(image)defon_click(event,x,y,flags,param):ifevent==cv2.EVENT_LBUTTONDOWN:# Foreground pointmasks,scores,_=predictor.predict(point_coords=np.array([[x,y]]),point_labels=np.array([1]),multimask_output=True)# Display best maskdisplay_mask(masks[np.argmax(scores)])
Workflow 2: Object extraction
defextract_object(image,point):"""Extract object at point with transparent background."""predictor.set_image(image)masks,scores,_=predictor.predict(point_coords=np.array([point]),point_labels=np.array([1]),multimask_output=True)best_mask=masks[np.argmax(scores)]# Create RGBA outputrgba=np.zeros((image.shape[0],image.shape[1],4),dtype=np.uint8)rgba[:,:,:3]=imagergba[:,:,3]=best_mask*255returnrgba
Workflow 3: Medical image segmentation
# Process medical images (grayscale to RGB)medical_image=cv2.imread("scan.png",cv2.IMREAD_GRAYSCALE)rgb_image=cv2.cvtColor(medical_image,cv2.COLOR_GRAY2RGB)predictor.set_image(rgb_image)# Segment region of interestmasks,scores,_=predictor.predict(box=np.array([x1,y1,x2,y2]),# ROI bounding boxmultimask_output=True)
frompycocotoolsimportmaskasmask_utils# Encode mask to RLErle=mask_utils.encode(np.asfortranarray(mask.astype(np.uint8)))rle["counts"]=rle["counts"].decode("utf-8")# Decode RLE to maskdecoded_mask=mask_utils.decode(rle)
Performance optimization
GPU memory
# Use smaller model for limited VRAMsam=sam_model_registry["vit_b"](checkpoint="sam_vit_b_01ec64.pth")# Process images in batches# Clear CUDA cache between large batchestorch.cuda.empty_cache()
Speed optimization
# Use half precisionsam=sam.half()# Reduce points for automatic generationmask_generator=SamAutomaticMaskGenerator(model=sam,points_per_side=16,# Default is 32)# Use ONNX for deployment# Export with --return-single-mask for faster inference