CoolFace
Apppublic

yuchen0187/Point-SAM

sourceHugging Facemitupdated 2y agoView on Hugging Face
11likes
app.py337 linesDownload Raw Back to root
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