Views
No views yet

1from PIL import Image
2import requests
3from transformers import SamHQModel, SamHQProcessor
4
5model = SamHQModel.from_pretrained("syscv-community/sam-hq-vit-base")
6processor = SamHQProcessor.from_pretrained("syscv-community/sam-hq-vit-base")
7
8img_url = "https://raw.githubusercontent.com/SysCV/sam-hq/refs/heads/main/demo/input_imgs/example1.png"
9raw_image = Image.open(requests.get(img_url, stream=True).raw).convert("RGB")
10input_boxes = [[[306, 132, 925, 893]]] # Bounding box for the image1inputs = processor(raw_image, input_boxes=input_boxes, return_tensors="pt").to("cuda")
2outputs = model(**inputs)
3masks = processor.image_processor.post_process_masks(outputs.pred_masks.cpu(), inputs["original_sizes"].cpu(), inputs["reshaped_input_sizes"].cpu())
4scores = outputs.iou_scores1024 points which are all fed to the model.points_per_batch argument)1from transformers import pipeline
2generator = pipeline("mask-generation", model="syscv-community/sam-hq-vit-base", device=0, points_per_batch=256)
3image_url = "https://raw.githubusercontent.com/SysCV/sam-hq/refs/heads/main/demo/input_imgs/example1.png"
4outputs = generator(image_url, points_per_batch=256)1import matplotlib.pyplot as plt
2from PIL import Image
3import numpy as np
4
5def show_mask(mask, ax, random_color=False):
6 if random_color:
7 color = np.concatenate([np.random.random(3), np.array([0.6])], axis=0)
8 else:
9 color = np.array([30 / 255, 144 / 255, 255 / 255, 0.6])
10 h, w = mask.shape[-2:]
11 mask_image = mask.reshape(h, w, 1) * color.reshape(1, 1, -1)
12 ax.imshow(mask_image)
13
14plt.imshow(np.array(raw_image))
15ax = plt.gca()
16for mask in outputs["masks"]:
17 show_mask(mask, ax=ax, random_color=True)
18plt.axis("off")
19plt.show()1import numpy as np
2import matplotlib.pyplot as plt
3def show_mask(mask, ax, random_color=False):
4 if random_color:
5 color = np.concatenate([np.random.random(3), np.array([0.6])], axis=0)
6 else:
7 color = np.array([30/255, 144/255, 255/255, 0.6])
8 h, w = mask.shape[-2:]
9 mask_image = mask.reshape(h, w, 1) * color.reshape(1, 1, -1)
10 ax.imshow(mask_image)
11def show_box(box, ax):
12 x0, y0 = box[0], box[1]
13 w, h = box[2] - box[0], box[3] - box[1]
14 ax.add_patch(plt.Rectangle((x0, y0), w, h, edgecolor='green', facecolor=(0,0,0,0), lw=2))
15def show_boxes_on_image(raw_image, boxes):
16 plt.figure(figsize=(10,10))
17 plt.imshow(raw_image)
18 for box in boxes:
19 show_box(box, plt.gca())
20 plt.axis('on')
21 plt.show()
22def show_points_on_image(raw_image, input_points, input_labels=None):
23 plt.figure(figsize=(10,10))
24 plt.imshow(raw_image)
25 input_points = np.array(input_points)
26 if input_labels is None:
27 labels = np.ones_like(input_points[:, 0])
28 else:
29 labels = np.array(input_labels)
30 show_points(input_points, labels, plt.gca())
31 plt.axis('on')
32 plt.show()
33def show_points_and_boxes_on_image(raw_image, boxes, input_points, input_labels=None):
34 plt.figure(figsize=(10,10))
35 plt.imshow(raw_image)
36 input_points = np.array(input_points)
37 if input_labels is None:
38 labels = np.ones_like(input_points[:, 0])
39 else:
40 labels = np.array(input_labels)
41 show_points(input_points, labels, plt.gca())
42 for box in boxes:
43 show_box(box, plt.gca())
44 plt.axis('on')
45 plt.show()
46def show_points_and_boxes_on_image(raw_image, boxes, input_points, input_labels=None):
47 plt.figure(figsize=(10,10))
48 plt.imshow(raw_image)
49 input_points = np.array(input_points)
50 if input_labels is None:
51 labels = np.ones_like(input_points[:, 0])
52 else:
53 labels = np.array(input_labels)
54 show_points(input_points, labels, plt.gca())
55 for box in boxes:
56 show_box(box, plt.gca())
57 plt.axis('on')
58 plt.show()
59def show_points(coords, labels, ax, marker_size=375):
60 pos_points = coords[labels==1]
61 neg_points = coords[labels==0]
62 ax.scatter(pos_points[:, 0], pos_points[:, 1], color='green', marker='*', s=marker_size, edgecolor='white', linewidth=1.25)
63 ax.scatter(neg_points[:, 0], neg_points[:, 1], color='red', marker='*', s=marker_size, edgecolor='white', linewidth=1.25)
64def show_masks_on_image(raw_image, masks, scores):
65 if len(masks.shape) == 4:
66 masks = masks.squeeze()
67 if scores.shape[0] == 1:
68 scores = scores.squeeze()
69 nb_predictions = scores.shape[-1]
70 fig, axes = plt.subplots(1, nb_predictions, figsize=(15, 15))
71 for i, (mask, score) in enumerate(zip(masks, scores)):
72 mask = mask.cpu().detach()
73 axes[i].imshow(np.array(raw_image))
74 show_mask(mask, axes[i])
75 axes[i].title.set_text(f"Mask {i+1}, Score: {score.item():.3f}")
76 axes[i].axis("off")
77 plt.show()
78def show_masks_on_single_image(raw_image, masks, scores):
79 if len(masks.shape) == 4:
80 masks = masks.squeeze()
81 if scores.shape[0] == 1:
82 scores = scores.squeeze()
83 # Convert image to numpy array if it's not already
84 image_np = np.array(raw_image)
85 # Create a figure
86 fig, ax = plt.subplots(figsize=(8, 8))
87 ax.imshow(image_np)
88 # Overlay all masks on the same image
89 for i, (mask, score) in enumerate(zip(masks, scores)):
90 mask = mask.cpu().detach().numpy() # Convert to NumPy
91 show_mask(mask, ax) # Assuming `show_mask` properly overlays the mask
92 ax.set_title(f"Overlayed Masks with Scores")
93 ax.axis("off")
94 plt.show()
95
96import torch
97from transformers import SamHQModel, SamHQProcessor
98
99device = "cuda" if torch.cuda.is_available() else "cpu"
100model = SamHQModel.from_pretrained("syscv-community/sam-hq-vit-base").to(device)
101processor = SamHQProcessor.from_pretrained("syscv-community/sam-hq-vit-base")
102
103from PIL import Image
104import requests
105img_url = "https://raw.githubusercontent.com/SysCV/sam-hq/refs/heads/main/demo/input_imgs/example1.png"
106raw_image = Image.open(requests.get(img_url, stream=True).raw).convert("RGB")
107plt.imshow(raw_image)
108
109inputs = processor(raw_image, return_tensors="pt").to(device)
110image_embeddings, intermediate_embeddings = model.get_image_embeddings(inputs["pixel_values"])
111
112input_boxes = [[[306, 132, 925, 893]]]
113show_boxes_on_image(raw_image, input_boxes[0])
114
115inputs.pop("pixel_values", None)
116inputs.update({"image_embeddings": image_embeddings})
117inputs.update({"intermediate_embeddings": intermediate_embeddings})
118with torch.no_grad():
119 outputs = model(**inputs)
120masks = processor.image_processor.post_process_masks(outputs.pred_masks.cpu(), inputs["original_sizes"].cpu(), inputs["reshaped_input_sizes"].cpu())
121scores = outputs.iou_scores
122
123show_masks_on_single_image(raw_image, masks[0], scores)
124
125show_masks_on_image(raw_image, masks[0], scores)@misc{ke2023segmenthighquality,
title={Segment Anything in High Quality},
author={Lei Ke and Mingqiao Ye and Martin Danelljan and Yifan Liu and Yu-Wing Tai and Chi-Keung Tang and Fisher Yu},
year={2023},
eprint={2306.01567},
archivePrefix={arXiv},
primaryClass={cs.CV},
url={https://arxiv.org/abs/2306.01567},
}