yuchen0187/Point-SAM
11
1import dataclasses2import os3 4import hydra5import numpy as np6import torch7from flask import Flask, jsonify, request, render_template8from flask_cors import CORS9from omegaconf import OmegaConf10from safetensors.torch import load_model11from scipy.spatial.transform import Rotation12 13from point_sam import build_point_sam14import argparse15 16app = Flask(__name__, static_folder="static")17CORS(app)18 19MAX_POINT_ID = 10020point_info_id = 021point_info_list = [None for _ in range(MAX_POINT_ID)]22 23@dataclasses.dataclass24class AuxInputs:25 coords: torch.Tensor26 features: torch.Tensor27 centers: torch.Tensor28 interp_index: torch.Tensor = None29 interp_weight: torch.Tensor = None30 31def repeat_interleave(x: torch.Tensor, repeats: int, dim: int):32 if repeats == 1:33 return x34 shape = list(x.shape)35 shape.insert(dim + 1, 1)36 shape[dim + 1] = repeats37 x = x.unsqueeze(dim + 1).expand(shape).flatten(dim, dim + 1)38 return x39 40 41class PointCloudProcessor:42 def __init__(self, device="cuda", batch=True, return_tensors="pt"):43 self.device = device44 self.batch = batch45 self.return_tensors = return_tensors46 47 self.center = None48 self.scale = None49 50 def __call__(self, xyz: np.ndarray, rgb: np.ndarray):51 # # The original data is z-up. Make it y-up.52 # rot = Rotation.from_euler("x", -90, degrees=True)53 # xyz = rot.apply(xyz)54 55 if self.center is None or self.scale is None:56 self.center = xyz.mean(0)57 self.scale = np.max(np.linalg.norm(xyz - self.center, axis=-1))58 59 xyz = (xyz - self.center) / self.scale60 rgb = ((rgb / 255.0) - 0.5) * 261 62 if self.return_tensors == "np":63 coords = np.float32(xyz)64 feats = np.float32(rgb)65 if self.batch:66 coords = np.expand_dims(coords, 0)67 feats = np.expand_dims(feats, 0)68 elif self.return_tensors == "pt":69 coords = torch.tensor(xyz, dtype=torch.float32, device=self.device)70 feats = torch.tensor(rgb, dtype=torch.float32, device=self.device)71 if self.batch:72 coords = coords.unsqueeze(0)73 feats = feats.unsqueeze(0)74 else:75 raise ValueError(self.return_tensors)76 77 return coords, feats78 79 def normalize(self, xyz):80 return (xyz - self.center) / self.scale81 82 83class PointCloudSAMPredictor:84 input_xyz: np.ndarray85 input_rgb: np.ndarray86 prompt_coords: list[tuple[float, float, float]]87 prompt_labels: list[int]88 89 coords: torch.Tensor90 feats: torch.Tensor91 92 pc_embedding: torch.Tensor93 patches: dict[str, torch.Tensor]94 prompt_mask: torch.Tensor95 96 def __init__(self):97 print("Created model")98 model = build_point_sam("./model-2.safetensors")99 model.pc_encoder.patch_embed.grouper.num_groups = 1024100 model.pc_encoder.patch_embed.grouper.group_size = 128101 if torch.cuda.is_available():102 model = model.cuda()103 model.eval()104 105 self.model = model106 107 self.input_rgb = None108 self.input_xyz = None109 110 self.input_processor = None111 self.coords = None112 self.feats = None113 114 self.pc_embedding = None115 self.patches = None116 117 self.prompt_coords = None118 self.prompt_labels = None119 self.prompt_mask = None120 self.candidate_index = 0121 122 @torch.no_grad()123 def set_pointcloud(self, xyz, rgb):124 self.input_xyz = xyz125 self.input_rgb = rgb126 127 self.input_processor = PointCloudProcessor()128 coords, feats = self.input_processor(xyz, rgb)129 self.coords = coords130 self.feats = feats131 132 pc_embedding, patches = self.model.pc_encoder(self.coords, self.feats)133 self.pc_embedding = pc_embedding134 self.patches = patches135 self.prompt_mask = None136 137 def set_prompts(self, prompt_coords, prompt_labels):138 self.prompt_coords = prompt_coords139 self.prompt_labels = prompt_labels140 141 @torch.no_grad()142 def predict_mask(self):143 normalized_prompt_coords = self.input_processor.normalize(144 np.array(self.prompt_coords)145 )146 prompt_coords = torch.tensor(147 normalized_prompt_coords, dtype=torch.float32, device="cuda"148 )149 prompt_labels = torch.tensor(150 self.prompt_labels, dtype=torch.bool, device="cuda"151 )152 prompt_coords = prompt_coords.reshape(1, -1, 3)153 prompt_labels = prompt_labels.reshape(1, -1)154 155 multimask_output = prompt_coords.shape[1] == 1156 157 # [B * M, num_outputs, num_points], [B * M, num_outputs]158 def decode_masks(coords, feats, pc_embedding, patches, prompt_coords, prompt_labels, prompt_masks, multimask_output):159 pc_embeddings, patches = pc_embedding, patches160 centers = patches["centers"]161 knn_idx = patches["knn_idx"]162 coords = patches["coords"]163 feats = patches["feats"]164 aux_inputs = AuxInputs(coords=coords, features=feats, centers=centers)165 166 pc_pe = self.model.point_encoder.pe_layer(centers)167 sparse_embeddings = self.model.point_encoder(prompt_coords, prompt_labels)168 dense_embeddings = self.model.mask_encoder(prompt_masks, coords, centers, knn_idx)169 dense_embeddings = repeat_interleave(170 dense_embeddings, sparse_embeddings.shape[0] // dense_embeddings.shape[0], 0171 )172 173 logits, iou_preds = self.model.mask_decoder(174 pc_embeddings,175 pc_pe,176 sparse_embeddings,177 dense_embeddings,178 aux_inputs=aux_inputs,179 multimask_output=multimask_output,180 )181 return logits, iou_preds182 183 logits, scores = decode_masks(184 self.coords,185 self.feats,186 self.pc_embedding,187 self.patches,188 prompt_coords,189 prompt_labels,190 self.prompt_mask[self.candidate_index].unsqueeze(0) if self.prompt_mask is not None else None,191 multimask_output,192 )193 logits = logits.squeeze(0)194 scores = scores.squeeze(0)195 196 # if multimask_output:197 # index = scores.argmax(0).item()198 # logit = logits[index]199 # else:200 # logit = logits.squeeze(0)201 202 # self.prompt_mask = logit.unsqueeze(0)203 204 # pred_mask = logit > 0205 # return pred_mask.cpu().numpy()206 207 # Sort according to scores208 _, indices = scores.sort(descending=True)209 logits = logits[indices]210 211 self.prompt_mask = logits # [num_outputs, num_points]212 self.candidate_index = 0213 214 return (logits > 0).cpu().numpy()215 216 def set_candidate(self, index):217 self.candidate_index = index218 219 220predictor = PointCloudSAMPredictor()221 222 223@app.route("/")224def index():225 return app.send_static_file("index.html")226 227@app.route("/assets/<path:path>")228def assets_route(path):229 print(path)230 return app.send_static_file(f"assets/{path}")231 232 233@app.route("/hello_world", methods=["GET"])234def hello_world():235 return "Hello, World!"236 237 238@app.route("/set_pointcloud", methods=["POST"])239def set_pointcloud():240 request_data = request.get_json()241 # print(request_data)242 # print(type(request_data["points"]))243 # print(type(request_data["colors"]))244 245 xyz = request_data["points"]246 xyz = np.array(xyz).reshape(-1, 3)247 rgb = request_data["colors"]248 rgb = np.array(list(rgb)).reshape(-1, 3)249 predictor.set_pointcloud(xyz, rgb)250 251 pc_embedding = predictor.pc_embedding.cpu()252 patches = {"centers": predictor.patches["centers"].cpu(), "knn_idx": predictor.patches["knn_idx"].cpu(), "coords": predictor.coords.cpu(), "feats": predictor.feats.cpu()}253 center = predictor.input_processor.center254 scale = predictor.input_processor.scale255 256 global point_info_id257 global point_info_list258 point_info_list[point_info_id] = {"pc_embedding": pc_embedding, "patches": patches, "center": center, "scale": scale, "prompt_mask": None}259 260 return_msg = {"user_id": point_info_id}261 point_info_id += 1262 return jsonify(return_msg)263 264 265@app.route("/set_candidate", methods=["POST"])266def set_candidate():267 request_data = request.get_json()268 candidate_index = request_data["index"]269 predictor.set_candidate(candidate_index)270 return "success"271 272 273def visualize_pcd_with_prompts(xyz, rgb, prompt_coords, prompt_labels):274 import trimesh275 276 pcd = trimesh.PointCloud(xyz, rgb)277 prompt_spheres = []278 for i, coord in enumerate(prompt_coords):279 sphere = trimesh.creation.icosphere()280 sphere.apply_scale(0.02)281 sphere.apply_translation(coord)282 sphere.visual.vertex_colors = [255, 0, 0] if prompt_labels[i] else [0, 255, 0]283 prompt_spheres.append(sphere)284 285 return trimesh.Scene([pcd] + prompt_spheres)286 287 288@app.route("/set_prompts", methods=["POST"])289def set_prompts():290 global point_info_list291 292 request_data = request.get_json()293 print(request_data.keys())294 295 # [n_prompts, 3]296 prompt_coords = request_data["prompt_coords"]297 # [n_prompts]. 0 for negative, 1 for positive298 prompt_labels = request_data["prompt_labels"]299 user_id = request_data["user_id"]300 print(user_id)301 point_info = point_info_list[user_id]302 predictor.pc_embedding = point_info["pc_embedding"].cuda()303 patches = point_info["patches"]304 predictor.patches = {"centers": patches["centers"].cuda(), "knn_idx": patches["knn_idx"].cuda(), "coords": patches["coords"].cuda(), "feats": patches["feats"].cuda()}305 predictor.input_processor.center = point_info["center"]306 predictor.input_processor.scale = point_info["scale"]307 if point_info["prompt_mask"] is not None:308 predictor.prompt_mask = point_info["prompt_mask"].cuda()309 else:310 predictor.prompt_mask = None311 # instance_id = request_data["instance_id"] # int312 if len(prompt_coords) == 0:313 predictor.prompt_mask = None314 pred_mask = np.zeros([len(prompt_coords)], dtype=np.bool_)315 return jsonify({"mask": pred_mask.tolist()})316 317 predictor.set_prompts(prompt_coords, prompt_labels)318 pred_mask = predictor.predict_mask()319 point_info_list[user_id]["prompt_mask"] = predictor.prompt_mask.cpu()320 321 # # Visualize322 # xyz = predictor.coords.cpu().numpy()[0]323 # rgb = predictor.feats.cpu().numpy()[0] * 0.5 + 0.5324 # prompt_coords = predictor.input_processor.normalize(np.array(predictor.prompt_coords))325 # scene = visualize_pcd_with_prompts(xyz, rgb, prompt_coords, predictor.prompt_labels)326 # scene.show()327 328 return jsonify({"mask": pred_mask.tolist()})329 330 331if __name__ == "__main__":332 parser = argparse.ArgumentParser()333 parser.add_argument("--host", type=str, default="0.0.0.0")334 parser.add_argument("--port", type=int, default=7860)335 args = parser.parse_args()336 app.run(host=args.host, port=args.port, debug=True)337 