Compare commits
1
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
cc4caff286 |
@@ -1,17 +1,14 @@
|
||||
#nodes.py
|
||||
|
||||
import torch
|
||||
import os
|
||||
import requests
|
||||
import numpy as np
|
||||
import logging
|
||||
import json
|
||||
import ast
|
||||
|
||||
import sys
|
||||
|
||||
import numpy as np #ugh
|
||||
import cv2 #double ugh
|
||||
|
||||
import pycuda.driver as cuda
|
||||
import pycuda.autoinit # Automatically initializes a CUDA context
|
||||
import cv2
|
||||
|
||||
# Add the directory containing 'sam2_realtime' to sys.path
|
||||
current_directory = os.path.dirname(os.path.abspath(__file__))
|
||||
@@ -19,7 +16,7 @@ sam2_realtime_path = os.path.join(current_directory) # Adjust the relative path
|
||||
sys.path.append(sam2_realtime_path)
|
||||
|
||||
from sam2_realtime.sam2_tensor_predictor import SAM2TensorPredictor
|
||||
from sam2_realtime.sam2_tensorrt_predictor import Sam2TensorrtPredictor
|
||||
from comfy.utils import load_torch_file
|
||||
|
||||
from omegaconf import OmegaConf
|
||||
from hydra.utils import instantiate
|
||||
@@ -36,7 +33,7 @@ class DownloadAndLoadSAM2RealtimeModel:
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"model": ([
|
||||
'sam2_hiera_tiny.pt', 'sam2_hiera_small.pt',
|
||||
'sam2_hiera_tiny.pt',
|
||||
],),
|
||||
"segmentor": (
|
||||
['realtime'],
|
||||
@@ -69,14 +66,12 @@ class DownloadAndLoadSAM2RealtimeModel:
|
||||
|
||||
download_path = os.path.join(folder_paths.models_dir, "sam2")
|
||||
model_path = os.path.join(download_path, model)
|
||||
print("model_path: ", model_path)
|
||||
|
||||
if not os.path.exists(download_path):
|
||||
os.makedirs(download_path)
|
||||
url = "https://dl.fbaipublicfiles.com/segment_anything_2/072824/sam2_hiera_tiny.pt"
|
||||
|
||||
if not os.path.exists(model_path):
|
||||
print(f"Downloading SAM2 model to: {model_path}")
|
||||
base_url = "https://dl.fbaipublicfiles.com/segment_anything_2/072824/"
|
||||
url = f"{base_url}{model}"
|
||||
response = requests.get(url, stream=True)
|
||||
response.raise_for_status()
|
||||
|
||||
@@ -89,9 +84,8 @@ class DownloadAndLoadSAM2RealtimeModel:
|
||||
|
||||
config_dir = os.path.join(script_directory, "sam2_configs")
|
||||
|
||||
model_cfg = model.replace(".pt", ".yaml")
|
||||
|
||||
# Code ripped out of sam2.build_sam.build_sam2_camera_predictor to appease Hydra
|
||||
model_cfg = "sam2_hiera_t.yaml" #TODO(pschroedl): remove hardcoded config and path
|
||||
with initialize_config_dir(config_dir=config_dir, version_base=None):
|
||||
cfg = compose(config_name=model_cfg)
|
||||
|
||||
@@ -146,194 +140,150 @@ class Sam2RealtimeSegmentation:
|
||||
return {
|
||||
"required": {
|
||||
"images": ("IMAGE",),
|
||||
# "sam2_model": ("SAM2MODEL",),
|
||||
"sam2_model": ("SAM2MODEL",),
|
||||
# "keep_model_loaded": ("BOOLEAN", {"default": True}),
|
||||
},
|
||||
"optional": {
|
||||
# "coordinates_positive": ("STRING", ),
|
||||
# "coordinates_negative": ("STRING", ),
|
||||
# "reset_tracking": ("BOOLEAN", {"default": False}),
|
||||
"optional": {
|
||||
"coordinates_positive": ("STRING", {"forceInput": True}),
|
||||
"point_labels": ("STRING", {"forceInput": True}),
|
||||
# "coordinates_negative": ("STRING", {"forceInput": True}),
|
||||
# "bboxes": ("BBOX", ),
|
||||
# "individual_objects": ("BOOLEAN", {"default": False}),
|
||||
# "mask": ("MASK", ),
|
||||
"threshold": ("FLOAT", {"forceInput": True}),
|
||||
"show_point": ("BOOLEAN", {"default": False}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_NAMES = ("PROCESSED_IMAGES", "MASK",)
|
||||
RETURN_TYPES = ("IMAGE", "MASK",)
|
||||
RETURN_NAMES = ("PROCESSED_IMAGES","MASK",)
|
||||
RETURN_TYPES = ("IMAGE", "IMAGE",)
|
||||
FUNCTION = "segment_images"
|
||||
CATEGORY = "SAM2-Realtime"
|
||||
|
||||
def __init__(self):
|
||||
##############################################################################
|
||||
# SETUP: Load Engines, Create Contexts, Allocate Buffers
|
||||
##############################################################################
|
||||
self.device = torch.device("cuda")
|
||||
|
||||
cuda.init()
|
||||
device = cuda.Device(0)
|
||||
context = device.make_context()
|
||||
print(f"Using device: {device.name()}")
|
||||
|
||||
# Allocate memory
|
||||
mem = cuda.mem_alloc(1024)
|
||||
print(f"Memory allocated at: {mem}")
|
||||
# context.pop()
|
||||
|
||||
tensor_models_path = os.path.join(folder_paths.models_dir, "tensorrt")
|
||||
self.predictor = Sam2TensorrtPredictor(
|
||||
encoder_engine_path=os.path.join(tensor_models_path, "sam2_hiera_tiny.encoder.engine"),
|
||||
decoder_engine_path=os.path.join(tensor_models_path, "sam2_hiera_tiny.decoder.engine")
|
||||
)
|
||||
self.predictor = None
|
||||
self.if_init = False
|
||||
|
||||
def _process_coordinate_input(self, coordinates, label):
|
||||
"""Helper function to process coordinate inputs safely"""
|
||||
if not coordinates:
|
||||
return [], []
|
||||
try:
|
||||
coord_list = ast.literal_eval(coordinates)
|
||||
points = [tuple(map(int, point)) for point in coord_list]
|
||||
labels = [label] * len(points)
|
||||
return points, labels
|
||||
except (ValueError, SyntaxError) as e:
|
||||
print(f"Error processing coordinates: {e}")
|
||||
return [], []
|
||||
|
||||
def _process_mask_logits(self, out_mask_logits, frame_shape, device):
|
||||
"""Helper function to process mask logits"""
|
||||
if out_mask_logits.shape[0] > 0:
|
||||
mask = (out_mask_logits[0, 0] > 0.5).byte()
|
||||
mask = torch.nn.functional.interpolate(
|
||||
mask.unsqueeze(0).unsqueeze(0).float(),
|
||||
size=frame_shape[:2],
|
||||
mode='nearest'
|
||||
).squeeze().byte().to(device)
|
||||
else:
|
||||
mask = torch.ones(frame_shape[:2], device=device, dtype=torch.uint8)
|
||||
return mask
|
||||
def _process_mask(self, mask: np.ndarray, frame_shape: tuple) -> np.ndarray:
|
||||
if mask.shape[0] == 0:
|
||||
logging.warning("Empty mask received")
|
||||
return np.zeros((frame_shape[0], frame_shape[1]), dtype="uint8")
|
||||
|
||||
colors = [
|
||||
[255, 0, 255], # Purple
|
||||
[0, 255, 255], # Yellow
|
||||
[255, 255, 0], # Cyan
|
||||
[0, 255, 0], # Green
|
||||
[255, 0, 0], # Blue
|
||||
]
|
||||
|
||||
combined_colored_mask = np.zeros((frame_shape[0], frame_shape[1], 4), dtype="uint8")
|
||||
|
||||
for i in range(mask.shape[0]):
|
||||
current_mask = (mask[i, 0] > 0).cpu().numpy().astype("uint8") * 255
|
||||
if current_mask.shape[:2] != frame_shape[:2]:
|
||||
current_mask = cv2.resize(current_mask, (frame_shape[1], frame_shape[0]))
|
||||
|
||||
# Create BGRA mask with transparency
|
||||
colored_mask = np.zeros((frame_shape[0], frame_shape[1], 4), dtype="uint8")
|
||||
color = colors[i % len(colors)]
|
||||
colored_mask[current_mask > 0] = color + [128] # Add alpha value of 128
|
||||
|
||||
# Alpha blend with existing masks
|
||||
alpha = colored_mask[:, :, 3:4] / 255.0
|
||||
combined_colored_mask = (1 - alpha) * combined_colored_mask + alpha * colored_mask
|
||||
|
||||
# Convert back to BGR for display
|
||||
combined_colored_mask = combined_colored_mask[:, :, :3].astype("uint8")
|
||||
return combined_colored_mask
|
||||
|
||||
def segment_images(
|
||||
self,
|
||||
images,
|
||||
# sam2_model,
|
||||
# coordinates_positive=None,
|
||||
sam2_model,
|
||||
# keep_model_loaded,
|
||||
coordinates_positive=None,
|
||||
# coordinates_negative=None,
|
||||
# reset_tracking=False,
|
||||
point_labels=None,
|
||||
# bboxes=None,
|
||||
# individual_objects=False,
|
||||
# mask=None,
|
||||
threshold=0.5,
|
||||
show_point=False,
|
||||
):
|
||||
|
||||
model = sam2_model["model"]
|
||||
device = sam2_model["device"]
|
||||
|
||||
device = torch.device("cuda")
|
||||
model.to(device)
|
||||
|
||||
processed_frames = []
|
||||
mask_list = []
|
||||
|
||||
# if reset_tracking:
|
||||
# self.if_init = False
|
||||
# The `model` is equivalent to `predictor` returned by sam2.build_sam.build_sam2_camera_predictor
|
||||
if self.predictor is None:
|
||||
self.predictor = model
|
||||
|
||||
with torch.inference_mode(), torch.autocast("cuda", dtype=torch.float16):
|
||||
for frame_idx, frame in enumerate(images):
|
||||
|
||||
# Tensor -> numpy array
|
||||
frame_numpy = frame.numpy() # shape: already (H,W,C)!!
|
||||
|
||||
# Convert from RGB -> BGR
|
||||
frame_bgr = cv2.cvtColor(frame_numpy, cv2.COLOR_RGB2BGR) # shape: (H,W,3), float32 in [0..1]
|
||||
|
||||
# Scale to [0..255] and cast to uint8
|
||||
frame_bgr = (frame_bgr * 255).clip(0,255).astype(np.uint8)
|
||||
def process_frame(frame, frame_idx):
|
||||
with torch.inference_mode(), torch.autocast("cuda", dtype=torch.float16):
|
||||
frame = frame.to(device).float()
|
||||
|
||||
if not self.if_init:
|
||||
self.predictor.load_first_frame(frame)
|
||||
self.if_init = True
|
||||
|
||||
out_mask_logits = self.predictor.load_first_frame_and_prompt(
|
||||
frame_bgr,
|
||||
point_coord=(512, 512)
|
||||
)
|
||||
coordinates_positive_list = ast.literal_eval(coordinates_positive)
|
||||
point_labels_list = ast.literal_eval(point_labels)
|
||||
point_labels_list = list(map(int, point_labels_list))
|
||||
|
||||
for idx, point in enumerate(coordinates_positive_list):
|
||||
point_tuple = tuple(map(int, point))
|
||||
_, _, out_mask_logits = self.predictor.add_new_prompt(
|
||||
frame_idx=0,
|
||||
obj_id=idx + 1,
|
||||
points=[point_tuple],
|
||||
labels=[point_labels_list[idx]]
|
||||
)
|
||||
else:
|
||||
out_mask_logits = self.predictor.track(frame_bgr)
|
||||
out_obj_ids, out_mask_logits = self.predictor.track(frame)
|
||||
|
||||
# Process mask logits
|
||||
# mask = self._process_mask_logits(out_mask_logits, frame.shape, self.device)
|
||||
mask = out_mask_logits
|
||||
if out_mask_logits.shape[0] > 0:
|
||||
mask = (out_mask_logits[0, 0] > threshold).byte()
|
||||
mask = torch.nn.functional.interpolate(
|
||||
mask.unsqueeze(0).unsqueeze(0).float(),
|
||||
size=(frame.shape[0], frame.shape[1]),
|
||||
mode='nearest'
|
||||
).squeeze(0).squeeze(0).byte() # Move the interpolated mask to the correct device
|
||||
else:
|
||||
mask = torch.ones((frame.shape[0], frame.shape[1]), device=device, dtype=torch.uint8)
|
||||
|
||||
# # Create colored overlay for processed frames
|
||||
# mask_colored = torch.stack([mask] * 3, dim=2)
|
||||
automask_colored = self._process_mask(mask,frame.shape)
|
||||
|
||||
# overlayed_frame = torch.add(frame * 0.7, mask_colored * 0.3)
|
||||
|
||||
# Create colored overlay for processed frames
|
||||
mask_colored = np.stack([mask] * 3, axis=2) # Stack along the last dimension (HxWxC format)
|
||||
# Draw points on the mask
|
||||
if show_point:
|
||||
for point in coordinates_positive:
|
||||
cv2.circle(automask_colored, tuple(point), radius=5, color=(0, 0, 255), thickness=-1)
|
||||
|
||||
mask_colored_resized = cv2.resize(mask_colored, (frame.shape[1], frame.shape[0]))
|
||||
automasked_frame = torch.add(frame * 0.7, automask_colored * 0.3)
|
||||
processed_frames.append(automasked_frame)
|
||||
|
||||
# Overlay the frame and the colored mask
|
||||
overlaid_frame = frame * 0.7 + mask_colored_resized * 0.3
|
||||
# TODO: This "mask" should be 1 channel to be returned as MASK type
|
||||
constructed_mask = torch.add(frame * 0.1, mask * 0.9)
|
||||
mask_list.append(constructed_mask)
|
||||
|
||||
processed_frames.append(torch.tensor(overlaid_frame))
|
||||
mask_list.append(torch.tensor(mask))
|
||||
|
||||
# Stack masks and frames
|
||||
for frame_idx, img in enumerate(images):
|
||||
process_frame(img, frame_idx)
|
||||
|
||||
stacked_masks = torch.stack(mask_list, dim=0)
|
||||
stacked_frames = torch.stack(processed_frames, dim=0)
|
||||
|
||||
stacked_frames = torch.stack(processed_frames, dim=0)
|
||||
return (stacked_frames, stacked_masks)
|
||||
|
||||
|
||||
|
||||
|
||||
class Sam2RealtimeSegmentationTest:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"images": ("IMAGE",),
|
||||
# "sam2_model": ("SAM2MODEL",),
|
||||
# "keep_model_loaded": ("BOOLEAN", {"default": True}),
|
||||
},
|
||||
"optional": {
|
||||
# "coordinates_positive": ("STRING", ),
|
||||
# "coordinates_negative": ("STRING", ),
|
||||
# "reset_tracking": ("BOOLEAN", {"default": False}),
|
||||
# "bboxes": ("BBOX", ),
|
||||
# "individual_objects": ("BOOLEAN", {"default": False}),
|
||||
# "mask": ("MASK", ),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_NAMES = ("PROCESSED_IMAGES", "MASK",)
|
||||
RETURN_TYPES = ("IMAGE", "MASK",)
|
||||
FUNCTION = "segment_images"
|
||||
CATEGORY = "SAM2-Realtime"
|
||||
|
||||
def segment_images(
|
||||
self,
|
||||
images,
|
||||
# sam2_model,
|
||||
# coordinates_positive=None,
|
||||
# coordinates_negative=None,
|
||||
# reset_tracking=False,
|
||||
):
|
||||
processed_frames = []
|
||||
mask_list = []
|
||||
# processed_frames.append(images)
|
||||
# images_temp = images
|
||||
# coordinates_positive_temp = coordinates_positive.__str__
|
||||
# reset_tracking_temp = reset_tracking
|
||||
|
||||
# Stack masks and frames
|
||||
stacked_masks = torch.stack(mask_list, dim=0)
|
||||
stacked_frames = torch.stack(processed_frames, dim=0)
|
||||
|
||||
return (stacked_frames, stacked_masks)
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"DownloadAndLoadSAM2RealtimeModel": DownloadAndLoadSAM2RealtimeModel,
|
||||
"Sam2RealtimeSegmentation": Sam2RealtimeSegmentation,
|
||||
"Sam2RealtimeSegmentationTest": Sam2RealtimeSegmentationTest
|
||||
"Sam2RealtimeSegmentation": Sam2RealtimeSegmentation
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"DownloadAndLoadSAM2RealtimeModel": "(Down)Load sam2_realtime Model",
|
||||
"Sam2RealtimeSegmentation": "Sam2RealtimeSegmentation",
|
||||
"Sam2RealtimeSegmentationTest": "Sam2RealtimeSegmentationTest"
|
||||
"Sam2RealtimeSegmentation": "Sam2RealtimeSegmentation"
|
||||
}
|
||||
|
||||
+1
-11
@@ -1,16 +1,6 @@
|
||||
pyyaml
|
||||
ninja
|
||||
numpy>=1.24.4
|
||||
tqdm>=4.66.1
|
||||
hydra-core>=1.3.2
|
||||
iopath>=0.1.10
|
||||
pillow>=9.4.0
|
||||
opencv-python>=4.5.0
|
||||
pycuda>=2022.1
|
||||
tensorrt-cu12==10.7.0
|
||||
tensorrt-cu12-bindings==10.7.0
|
||||
tensorrt-cu12-libs==10.7.0
|
||||
tensorrt==10.7.0
|
||||
# git+https://github.com/pschroedl/ComfyUI-SAM2-Realtime.git@main#egg=sam2_realtime
|
||||
pip install -e .
|
||||
|
||||
pillow>=9.4.0
|
||||
@@ -1,116 +0,0 @@
|
||||
# @package _global_
|
||||
|
||||
# Model
|
||||
model:
|
||||
_target_: sam2_realtime.modeling.sam2_base.SAM2Base
|
||||
image_encoder:
|
||||
_target_: sam2_realtime.modeling.backbones.image_encoder.ImageEncoder
|
||||
scalp: 1
|
||||
trunk:
|
||||
_target_: sam2_realtime.modeling.backbones.hieradet.Hiera
|
||||
embed_dim: 96
|
||||
num_heads: 1
|
||||
stages: [1, 2, 11, 2]
|
||||
global_att_blocks: [7, 10, 13]
|
||||
window_pos_embed_bkg_spatial_size: [7, 7]
|
||||
neck:
|
||||
_target_: sam2_realtime.modeling.backbones.image_encoder.FpnNeck
|
||||
position_encoding:
|
||||
_target_: sam2_realtime.modeling.position_encoding.PositionEmbeddingSine
|
||||
num_pos_feats: 256
|
||||
normalize: true
|
||||
scale: null
|
||||
temperature: 10000
|
||||
d_model: 256
|
||||
backbone_channel_list: [768, 384, 192, 96]
|
||||
fpn_top_down_levels: [2, 3] # output level 0 and 1 directly use the backbone features
|
||||
fpn_interp_model: nearest
|
||||
|
||||
memory_attention:
|
||||
_target_: sam2_realtime.modeling.memory_attention.MemoryAttention
|
||||
d_model: 256
|
||||
pos_enc_at_input: true
|
||||
layer:
|
||||
_target_: sam2_realtime.modeling.memory_attention.MemoryAttentionLayer
|
||||
activation: relu
|
||||
dim_feedforward: 2048
|
||||
dropout: 0.1
|
||||
pos_enc_at_attn: false
|
||||
self_attention:
|
||||
_target_: sam2_realtime.modeling.sam.transformer.RoPEAttention
|
||||
rope_theta: 10000.0
|
||||
feat_sizes: [32, 32]
|
||||
embedding_dim: 256
|
||||
num_heads: 1
|
||||
downsample_rate: 1
|
||||
dropout: 0.1
|
||||
d_model: 256
|
||||
pos_enc_at_cross_attn_keys: true
|
||||
pos_enc_at_cross_attn_queries: false
|
||||
cross_attention:
|
||||
_target_: sam2_realtime.modeling.sam.transformer.RoPEAttention
|
||||
rope_theta: 10000.0
|
||||
feat_sizes: [32, 32]
|
||||
rope_k_repeat: True
|
||||
embedding_dim: 256
|
||||
num_heads: 1
|
||||
downsample_rate: 1
|
||||
dropout: 0.1
|
||||
kv_in_dim: 64
|
||||
num_layers: 4
|
||||
|
||||
memory_encoder:
|
||||
_target_: sam2_realtime.modeling.memory_encoder.MemoryEncoder
|
||||
out_dim: 64
|
||||
position_encoding:
|
||||
_target_: sam2_realtime.modeling.position_encoding.PositionEmbeddingSine
|
||||
num_pos_feats: 64
|
||||
normalize: true
|
||||
scale: null
|
||||
temperature: 10000
|
||||
mask_downsampler:
|
||||
_target_: sam2_realtime.modeling.memory_encoder.MaskDownSampler
|
||||
kernel_size: 3
|
||||
stride: 2
|
||||
padding: 1
|
||||
fuser:
|
||||
_target_: sam2_realtime.modeling.memory_encoder.Fuser
|
||||
layer:
|
||||
_target_: sam2_realtime.modeling.memory_encoder.CXBlock
|
||||
dim: 256
|
||||
kernel_size: 7
|
||||
padding: 3
|
||||
layer_scale_init_value: 1e-6
|
||||
use_dwconv: True # depth-wise convs
|
||||
num_layers: 2
|
||||
|
||||
num_maskmem: 7
|
||||
image_size: 512
|
||||
# apply scaled sigmoid on mask logits for memory encoder, and directly feed input mask as output mask
|
||||
sigmoid_scale_for_mem_enc: 20.0
|
||||
sigmoid_bias_for_mem_enc: -10.0
|
||||
use_mask_input_as_output_without_sam: true
|
||||
# Memory
|
||||
directly_add_no_mem_embed: true
|
||||
# use high-resolution feature map in the SAM mask decoder
|
||||
use_high_res_features_in_sam: true
|
||||
# output 3 masks on the first click on initial conditioning frames
|
||||
multimask_output_in_sam: true
|
||||
# SAM heads
|
||||
iou_prediction_use_sigmoid: True
|
||||
# cross-attend to object pointers from other frames (based on SAM output tokens) in the encoder
|
||||
use_obj_ptrs_in_encoder: true
|
||||
add_tpos_enc_to_obj_ptrs: false
|
||||
only_obj_ptrs_in_the_past_for_eval: true
|
||||
# object occlusion prediction
|
||||
pred_obj_scores: true
|
||||
pred_obj_scores_mlp: true
|
||||
fixed_no_obj_ptr: true
|
||||
# multimask tracking settings
|
||||
multimask_output_for_tracking: true
|
||||
use_multimask_token_for_obj_ptr: true
|
||||
multimask_min_pt_num: 0
|
||||
multimask_max_pt_num: 1
|
||||
use_mlp_for_obj_ptr_proj: true
|
||||
# Compilation flag
|
||||
compile_image_encoder: False
|
||||
@@ -85,7 +85,7 @@ model:
|
||||
num_layers: 2
|
||||
|
||||
num_maskmem: 7
|
||||
image_size: 512
|
||||
image_size: 1024
|
||||
# apply scaled sigmoid on mask logits for memory encoder, and directly feed input mask as output mask
|
||||
# SAM decoder
|
||||
sigmoid_scale_for_mem_enc: 20.0
|
||||
@@ -14,6 +14,7 @@ from tqdm import tqdm
|
||||
from sam2_realtime.modeling.sam2_base import NO_OBJ_SCORE, SAM2Base
|
||||
from sam2_realtime.utils.misc import concat_points, fill_holes_in_mask_scores, load_video_frames
|
||||
|
||||
|
||||
class SAM2TensorPredictor(SAM2Base):
|
||||
"""The predictor class to handle user interactions and manage inference states."""
|
||||
|
||||
@@ -54,9 +55,7 @@ class SAM2TensorPredictor(SAM2Base):
|
||||
img = img.float()
|
||||
else:
|
||||
raise ValueError("Input must be a numpy array or a PyTorch tensor")
|
||||
#save original height/width
|
||||
orig_h, orig_w = img.shape[1:]
|
||||
|
||||
|
||||
# Resize to the target size (supports tensor resizing)
|
||||
img = torch.nn.functional.interpolate(
|
||||
img.unsqueeze(0), size=(image_size, image_size), mode="bilinear", align_corners=False
|
||||
@@ -69,29 +68,28 @@ class SAM2TensorPredictor(SAM2Base):
|
||||
img /= img_std
|
||||
|
||||
height, width = img.shape[1:] # CHW format
|
||||
return img, width, height, orig_w, orig_h
|
||||
return img, width, height
|
||||
|
||||
@torch.inference_mode()
|
||||
def load_first_frame(self, img):
|
||||
if isinstance(img, torch.Tensor):
|
||||
img = img.to(self.device) # Ensure the tensor is on the correct device
|
||||
|
||||
|
||||
self.condition_state = self._init_state(
|
||||
offload_video_to_cpu=False, offload_state_to_cpu=False
|
||||
)
|
||||
img, width, height, orig_w, orig_h = self.prepare_data(img, image_size=self.image_size)
|
||||
self._orig_hw = (orig_w, orig_h)
|
||||
img, width, height = self.prepare_data(img, image_size=self.image_size)
|
||||
self.condition_state["images"] = [img]
|
||||
self.condition_state["num_frames"] = len(self.condition_state["images"])
|
||||
self.condition_state["video_height"] = height
|
||||
self.condition_state["video_width"] = width
|
||||
self._get_image_feature(frame_idx=0, batch_size=1)
|
||||
|
||||
|
||||
def add_conditioning_frame(self, img):
|
||||
if isinstance(img, torch.Tensor):
|
||||
img = img.to(self.device) # Ensure the tensor is on the correct device
|
||||
|
||||
img, width, height, _, _ = self.prepare_data(img, image_size=self.image_size)
|
||||
img, width, height = self.prepare_data(img, image_size=self.image_size)
|
||||
self.condition_state["images"].append(img)
|
||||
self.condition_state["num_frames"] = len(self.condition_state["images"])
|
||||
self._get_image_feature(
|
||||
@@ -237,15 +235,14 @@ class SAM2TensorPredictor(SAM2Base):
|
||||
points = torch.cat([box_coords, points], dim=1)
|
||||
labels = torch.cat([box_labels, labels], dim=1)
|
||||
if normalize_coords:
|
||||
#video_H = self.condition_state["video_height"]
|
||||
#video_W = self.condition_state["video_width"]
|
||||
orig_w, orig_h = self._orig_hw
|
||||
|
||||
points = points / torch.tensor([orig_w, orig_h]).to(points.device)
|
||||
video_H = self.condition_state["video_height"]
|
||||
video_W = self.condition_state["video_width"]
|
||||
points = points / torch.tensor([video_W, video_H]).to(points.device)
|
||||
# scale the (normalized) coordinates by the model's internal image size
|
||||
points = points * self.image_size
|
||||
points = points.to(self.condition_state["device"])
|
||||
labels = labels.to(self.condition_state["device"])
|
||||
|
||||
if not clear_old_points:
|
||||
point_inputs = point_inputs_per_frame.get(frame_idx, None)
|
||||
else:
|
||||
@@ -345,16 +342,14 @@ class SAM2TensorPredictor(SAM2Base):
|
||||
if labels.dim() == 1:
|
||||
labels = labels.unsqueeze(0) # add batch dimension
|
||||
if normalize_coords:
|
||||
#video_H = self.condition_state["video_height"]
|
||||
#video_W = self.condition_state["video_width"]
|
||||
orig_w, orig_h = self._orig_hw
|
||||
|
||||
points = points / torch.tensor([orig_w, orig_h]).to(points.device)
|
||||
|
||||
video_H = self.condition_state["video_height"]
|
||||
video_W = self.condition_state["video_width"]
|
||||
points = points / torch.tensor([video_W, video_H]).to(points.device)
|
||||
# scale the (normalized) coordinates by the model's internal image size
|
||||
points = points * self.image_size
|
||||
points = points.to(self.condition_state["device"])
|
||||
labels = labels.to(self.condition_state["device"])
|
||||
|
||||
if not clear_old_points:
|
||||
point_inputs = point_inputs_per_frame.get(frame_idx, None)
|
||||
else:
|
||||
@@ -774,7 +769,7 @@ class SAM2TensorPredictor(SAM2Base):
|
||||
if isinstance(img, torch.Tensor):
|
||||
img = img.to(self.device) # Ensure the tensor is on the correct device
|
||||
|
||||
img, _, _ , _, _ = self.prepare_data(img, image_size=self.image_size)
|
||||
img, _, _ = self.prepare_data(img, image_size=self.image_size)
|
||||
|
||||
output_dict = self.condition_state["output_dict"]
|
||||
obj_ids = self.condition_state["obj_ids"]
|
||||
|
||||
@@ -1,241 +0,0 @@
|
||||
import sys
|
||||
import numpy as np
|
||||
import cv2
|
||||
|
||||
import pycuda.driver as cuda
|
||||
import pycuda.autoinit # Automatically initializes a CUDA context
|
||||
import tensorrt as trt
|
||||
|
||||
TRT_LOGGER = trt.Logger(trt.Logger.WARNING)
|
||||
|
||||
ENCODER_ENGINE_PATH = "sam2_hiera_tiny.encoder.engine"
|
||||
DECODER_ENGINE_PATH = "sam2_hiera_tiny.decoder.engine"
|
||||
|
||||
def load_engine(engine_path: str) -> trt.ICudaEngine:
|
||||
"""
|
||||
Load a TensorRT engine from a file.
|
||||
"""
|
||||
with open(engine_path, "rb") as f:
|
||||
engine_data = f.read()
|
||||
runtime = trt.Runtime(TRT_LOGGER)
|
||||
engine = runtime.deserialize_cuda_engine(engine_data)
|
||||
return engine
|
||||
|
||||
def allocate_io_tensors(engine: trt.ICudaEngine):
|
||||
"""
|
||||
Allocates device+host buffers for each I/O tensor in name-based I/O.
|
||||
Returns (device_buffers, host_buffers, input_names, output_names).
|
||||
"""
|
||||
device_buffers = {}
|
||||
host_buffers = {}
|
||||
input_names = []
|
||||
output_names = []
|
||||
|
||||
n_tensors = engine.num_io_tensors
|
||||
for i in range(n_tensors):
|
||||
tname = engine.get_tensor_name(i)
|
||||
mode = engine.get_tensor_mode(tname) # trt.TensorIOMode.INPUT or OUTPUT
|
||||
shape = engine.get_tensor_shape(tname)
|
||||
dtype = engine.get_tensor_dtype(tname)
|
||||
|
||||
# Convert TRT dtype -> numpy dtype
|
||||
if dtype == trt.float32:
|
||||
np_dtype = np.float32
|
||||
elif dtype == trt.float16:
|
||||
np_dtype = np.float16
|
||||
else:
|
||||
raise ValueError(f"Unsupported TRT dtype: {dtype}")
|
||||
|
||||
volume = np.prod(shape)
|
||||
host_mem = np.zeros(volume, dtype=np_dtype)
|
||||
device_mem = cuda.mem_alloc(host_mem.nbytes)
|
||||
|
||||
device_buffers[tname] = device_mem
|
||||
host_buffers[tname] = host_mem
|
||||
|
||||
if mode == trt.TensorIOMode.INPUT:
|
||||
input_names.append(tname)
|
||||
else:
|
||||
output_names.append(tname)
|
||||
|
||||
return device_buffers, host_buffers, input_names, output_names
|
||||
|
||||
def do_inference_async(
|
||||
context: trt.IExecutionContext,
|
||||
device_buffers: dict,
|
||||
host_buffers: dict,
|
||||
input_names: list,
|
||||
output_names: list,
|
||||
cuda_stream: cuda.Stream
|
||||
):
|
||||
"""
|
||||
1) Copy inputs (host->device)
|
||||
2) context.set_tensor_address(...)
|
||||
3) context.execute_async_v3(cuda_stream.handle)
|
||||
4) stream.synchronize()
|
||||
5) Copy outputs (device->host)
|
||||
6) return output in the same order as output_names
|
||||
"""
|
||||
# 1) Copy inputs to device
|
||||
for tname in input_names:
|
||||
cuda.memcpy_htod_async(device_buffers[tname], host_buffers[tname], stream=cuda_stream)
|
||||
|
||||
# 2) Set addresses in context
|
||||
for tname in input_names + output_names:
|
||||
context.set_tensor_address(tname, int(device_buffers[tname]))
|
||||
|
||||
# 3) Execute
|
||||
success = context.execute_async_v3(cuda_stream.handle)
|
||||
if not success:
|
||||
raise RuntimeError("TensorRT execute_async_v3() returned False.")
|
||||
|
||||
# 4) Synchronize to ensure inference is done before copying back
|
||||
cuda_stream.synchronize()
|
||||
|
||||
# 5) Copy outputs back to host
|
||||
for tname in output_names:
|
||||
cuda.memcpy_dtoh_async(host_buffers[tname], device_buffers[tname], stream=cuda_stream)
|
||||
|
||||
# Wait until all D->H copies are done
|
||||
cuda_stream.synchronize()
|
||||
|
||||
# 6) Gather outputs
|
||||
outputs = [host_buffers[t].copy() for t in output_names]
|
||||
return outputs
|
||||
|
||||
def run_encoder_async(
|
||||
image_bgr: np.ndarray,
|
||||
context: trt.IExecutionContext,
|
||||
device_buffers: dict,
|
||||
host_buffers: dict,
|
||||
input_names: list,
|
||||
output_names: list,
|
||||
image_size: tuple = (1024, 1024),
|
||||
cuda_stream: cuda.Stream = None
|
||||
):
|
||||
"""
|
||||
Run the SAM encoder in async mode.
|
||||
"""
|
||||
if cuda_stream is None:
|
||||
cuda_stream = cuda.Stream()
|
||||
|
||||
# Preprocess image
|
||||
H, W = image_size
|
||||
resized = cv2.resize(image_bgr, (W, H))
|
||||
rgb = cv2.cvtColor(resized, cv2.COLOR_BGR2RGB).astype(np.float32) / 255.0
|
||||
nchw = np.transpose(rgb, (2, 0, 1))[None, ...] # (1,3,H,W)
|
||||
|
||||
if len(input_names) != 1:
|
||||
raise RuntimeError(f"Expected exactly 1 input for encoder, got {input_names}")
|
||||
|
||||
in_name = input_names[0]
|
||||
host_buffers[in_name][:] = nchw.flatten()
|
||||
|
||||
# Async inference
|
||||
outputs = do_inference_async(
|
||||
context,
|
||||
device_buffers,
|
||||
host_buffers,
|
||||
input_names,
|
||||
output_names,
|
||||
cuda_stream
|
||||
)
|
||||
|
||||
# Expect 3 outputs
|
||||
if len(outputs) != 3:
|
||||
raise RuntimeError(f"Encoder expected 3 outputs, got {len(outputs)}")
|
||||
|
||||
hr0_flat, hr1_flat, emb_flat = outputs
|
||||
# Reshape according to model's output shapes
|
||||
# Shapes for sam2_hiera_tiny:
|
||||
# hr0: (1, 32, 256, 256)
|
||||
# hr1: (1, 64, 128, 128)
|
||||
# emb: (1, 256, 64, 64)
|
||||
hr0 = hr0_flat.reshape((1, 32, 256, 256))
|
||||
hr1 = hr1_flat.reshape((1, 64, 128, 128))
|
||||
emb = emb_flat.reshape((1, 256, 64, 64))
|
||||
|
||||
print("[DEBUG] Encoder outputs shapes:", hr0.shape, hr1.shape, emb.shape)
|
||||
return hr0, hr1, emb
|
||||
|
||||
def run_decoder_async(
|
||||
high_res_feats_0: np.ndarray,
|
||||
high_res_feats_1: np.ndarray,
|
||||
image_embed: np.ndarray,
|
||||
point_coords: np.ndarray,
|
||||
point_labels: np.ndarray,
|
||||
mask_input: np.ndarray,
|
||||
has_mask_input: np.ndarray,
|
||||
context: trt.IExecutionContext,
|
||||
device_buffers: dict,
|
||||
host_buffers: dict,
|
||||
input_names: list,
|
||||
output_names: list,
|
||||
cuda_stream: cuda.Stream = None
|
||||
):
|
||||
"""
|
||||
Run the SAM decoder in async mode with 7 inputs:
|
||||
0) image_embed
|
||||
1) high_res_feats_0
|
||||
2) high_res_feats_1
|
||||
3) point_coords
|
||||
4) point_labels
|
||||
5) mask_input
|
||||
6) has_mask_input
|
||||
"""
|
||||
if cuda_stream is None:
|
||||
cuda_stream = cuda.Stream()
|
||||
|
||||
if len(input_names) != 7:
|
||||
raise RuntimeError(f"Decoder expects 7 inputs, got {len(input_names)}: {input_names}")
|
||||
|
||||
# Flatten & place data
|
||||
host_buffers[input_names[0]][:] = image_embed.flatten()
|
||||
host_buffers[input_names[1]][:] = high_res_feats_0.flatten()
|
||||
host_buffers[input_names[2]][:] = high_res_feats_1.flatten()
|
||||
host_buffers[input_names[3]][:] = point_coords.flatten()
|
||||
host_buffers[input_names[4]][:] = point_labels.flatten()
|
||||
host_buffers[input_names[5]][:] = mask_input.flatten()
|
||||
host_buffers[input_names[6]][:] = has_mask_input.flatten()
|
||||
|
||||
outputs = do_inference_async(
|
||||
context,
|
||||
device_buffers,
|
||||
host_buffers,
|
||||
input_names,
|
||||
output_names,
|
||||
cuda_stream
|
||||
)
|
||||
|
||||
# Typically 2 outputs: masks, iou_predictions
|
||||
if len(outputs) != 2:
|
||||
raise RuntimeError(f"Decoder expected 2 outputs, got {len(outputs)}")
|
||||
|
||||
masks_flat, iou_flat = outputs
|
||||
# Reshape as needed for your model
|
||||
# Example: (1,3,256,256) for masks and (1,3) for iou_preds
|
||||
masks = masks_flat.reshape((1, 3, 256, 256))
|
||||
iou_preds = iou_flat.reshape((1, 3))
|
||||
print("Decoder outputs shapes:", masks.shape, iou_preds.shape)
|
||||
return masks, iou_preds
|
||||
|
||||
def prepare_prompts(
|
||||
batch_size: int = 1,
|
||||
image_size: tuple = (1024, 1024),
|
||||
):
|
||||
"""
|
||||
Create minimal prompt inputs for the decoder:
|
||||
- point_coords (B, 1, 2)
|
||||
- point_labels (B, 1)
|
||||
- mask_input (B, 1, 256, 256) (for sam2_hiera_tiny)
|
||||
- has_mask_input (B,)
|
||||
"""
|
||||
H, W = image_size
|
||||
point_coords = np.array([[[W // 2, H // 2]]], dtype=np.float32) # shape (B,1,2)
|
||||
point_labels = np.array([[1]], dtype=np.float32) # shape (B,1)
|
||||
|
||||
# For sam2_hiera_tiny => final embedding is 64×64, so mask_input is 4× that => 256×256
|
||||
mask_input = np.zeros((batch_size, 1, 256, 256), dtype=np.float32)
|
||||
has_mask_input = np.zeros((batch_size,), dtype=np.float32)
|
||||
|
||||
return point_coords, point_labels, mask_input, has_mask_input
|
||||
@@ -1,203 +0,0 @@
|
||||
import numpy as np
|
||||
import cv2
|
||||
|
||||
from sam2_realtime.sam2_tensorrt import (
|
||||
load_engine,
|
||||
allocate_io_tensors,
|
||||
run_encoder_async,
|
||||
run_decoder_async,
|
||||
prepare_prompts
|
||||
)
|
||||
|
||||
class Sam2TensorrtPredictor:
|
||||
"""
|
||||
A simplified single-object predictor that:
|
||||
- Uses the TensorRT-based SAM2 encoder/decoder.
|
||||
- Tracks across frames by feeding the previously predicted mask as a prompt.
|
||||
- Allows resetting the tracking with a new point prompt (i.e., ignoring prior mask).
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
encoder_engine_path: str,
|
||||
decoder_engine_path: str,
|
||||
image_size=(1024, 1024)
|
||||
):
|
||||
"""
|
||||
1) Load TensorRT engines
|
||||
2) Create execution contexts
|
||||
3) Allocate device+host buffers
|
||||
"""
|
||||
self.image_size = image_size # (height, width)
|
||||
|
||||
print("[Sam2TensorrtPredictor] Loading TensorRT engines...")
|
||||
self.encoder_engine = load_engine(encoder_engine_path)
|
||||
self.decoder_engine = load_engine(decoder_engine_path)
|
||||
|
||||
# Create execution contexts
|
||||
self.encoder_context = self.encoder_engine.create_execution_context()
|
||||
self.decoder_context = self.decoder_engine.create_execution_context()
|
||||
|
||||
# Allocate buffers for encoder
|
||||
(
|
||||
self.enc_device_bufs,
|
||||
self.enc_host_bufs,
|
||||
self.enc_in_names,
|
||||
self.enc_out_names
|
||||
) = allocate_io_tensors(self.encoder_engine)
|
||||
|
||||
# Allocate buffers for decoder
|
||||
(
|
||||
self.dec_device_bufs,
|
||||
self.dec_host_bufs,
|
||||
self.dec_in_names,
|
||||
self.dec_out_names
|
||||
) = allocate_io_tensors(self.decoder_engine)
|
||||
|
||||
print("[Sam2TensorrtPredictor] Engines loaded and buffers allocated.")
|
||||
|
||||
# Internal states
|
||||
self._initialized = False
|
||||
self.prev_mask = None
|
||||
self.prev_feats0 = None
|
||||
self.prev_feats1 = None
|
||||
self.prev_embed = None
|
||||
|
||||
def load_first_frame_and_prompt(self, frame_bgr: np.ndarray, point_coord, point_label=1):
|
||||
"""
|
||||
1) Encode the first frame to get feats0, feats1, embed.
|
||||
2) Use the user-provided point prompt to decode an initial mask.
|
||||
3) Store that mask in self.prev_mask (for subsequent frames).
|
||||
"""
|
||||
self._initialized = False
|
||||
self.prev_mask = None
|
||||
self.prev_feats0 = None
|
||||
self.prev_feats1 = None
|
||||
self.prev_embed = None
|
||||
|
||||
# -- (1) run the encoder
|
||||
feats0, feats1, embed = run_encoder_async(
|
||||
image_bgr=frame_bgr,
|
||||
context=self.encoder_context,
|
||||
device_buffers=self.enc_device_bufs,
|
||||
host_buffers=self.enc_host_bufs,
|
||||
input_names=self.enc_in_names,
|
||||
output_names=self.enc_out_names,
|
||||
image_size=self.image_size,
|
||||
)
|
||||
|
||||
# -- (2) build the prompt arrays
|
||||
batch_size = 1
|
||||
pt_coords, pt_labels, zero_mask, zero_has_mask = prepare_prompts(
|
||||
batch_size=batch_size,
|
||||
image_size=self.image_size
|
||||
)
|
||||
|
||||
# Overwrite the default center prompt with the user-provided point
|
||||
# Note: point_coord = (x, y)
|
||||
pt_coords[0, 0, 0] = float(point_coord[0]) # x
|
||||
pt_coords[0, 0, 1] = float(point_coord[1]) # y
|
||||
pt_labels[0, 0] = float(point_label)
|
||||
|
||||
# decode (prompt-based)
|
||||
masks, iou_preds = run_decoder_async(
|
||||
high_res_feats_0=feats0,
|
||||
high_res_feats_1=feats1,
|
||||
image_embed=embed,
|
||||
point_coords=pt_coords,
|
||||
point_labels=pt_labels,
|
||||
mask_input=zero_mask, # no prior mask for the first frame
|
||||
has_mask_input=zero_has_mask,
|
||||
context=self.decoder_context,
|
||||
device_buffers=self.dec_device_bufs,
|
||||
host_buffers=self.dec_host_bufs,
|
||||
input_names=self.dec_in_names,
|
||||
output_names=self.dec_out_names,
|
||||
)
|
||||
|
||||
# pick best mask
|
||||
best_idx = np.argmax(iou_preds[0]) # choose the highest IOU channel out of 3
|
||||
mask_logit = masks[0, best_idx] # shape: (256, 256)
|
||||
mask_prob = 1 / (1 + np.exp(-mask_logit))
|
||||
mask_bin = (mask_prob > 0.5).astype(np.uint8)
|
||||
|
||||
# -- (3) store for next step
|
||||
self.prev_mask = mask_bin
|
||||
self.prev_feats0 = feats0
|
||||
self.prev_feats1 = feats1
|
||||
self.prev_embed = embed
|
||||
self._initialized = True
|
||||
|
||||
print("[load_first_frame_and_prompt] First mask computed from point prompt.")
|
||||
return mask_bin
|
||||
|
||||
def track(self, frame_bgr: np.ndarray):
|
||||
"""
|
||||
For each subsequent frame:
|
||||
1) Encode the new frame
|
||||
2) Provide the previously predicted mask as 'mask_input' prompt
|
||||
3) Pick best mask => new self.prev_mask
|
||||
"""
|
||||
if not self._initialized or self.prev_mask is None:
|
||||
raise RuntimeError(
|
||||
"Sam2TensorrtPredictor not initialized with a first frame/prompt. "
|
||||
"Call load_first_frame_and_prompt(...) first."
|
||||
)
|
||||
|
||||
# 1) Encode the new frame
|
||||
feats0, feats1, embed = run_encoder_async(
|
||||
image_bgr=frame_bgr,
|
||||
context=self.encoder_context,
|
||||
device_buffers=self.enc_device_bufs,
|
||||
host_buffers=self.enc_host_bufs,
|
||||
input_names=self.enc_in_names,
|
||||
output_names=self.enc_out_names,
|
||||
image_size=self.image_size,
|
||||
)
|
||||
|
||||
# 2) Prepare mask_input from the previously predicted binary mask
|
||||
if self.prev_mask.shape != (256, 256):
|
||||
# if stored differently, resize to 256x256
|
||||
pm = cv2.resize(
|
||||
self.prev_mask.astype(np.float32),
|
||||
(256, 256),
|
||||
interpolation=cv2.INTER_LINEAR
|
||||
)[None, None] # shape: (1,1,256,256)
|
||||
else:
|
||||
pm = self.prev_mask[None, None].astype(np.float32)
|
||||
|
||||
has_mask_input = np.array([1.0], dtype=np.float32) # shape: (1,)
|
||||
|
||||
# 3) We want no new user points, we set them to dummy zero
|
||||
point_coords = np.zeros((1, 1, 2), dtype=np.float32)
|
||||
point_labels = np.zeros((1, 1), dtype=np.float32)
|
||||
|
||||
masks, iou_preds = run_decoder_async(
|
||||
high_res_feats_0=feats0,
|
||||
high_res_feats_1=feats1,
|
||||
image_embed=embed,
|
||||
point_coords=point_coords,
|
||||
point_labels=point_labels,
|
||||
mask_input=pm,
|
||||
has_mask_input=has_mask_input,
|
||||
context=self.decoder_context,
|
||||
device_buffers=self.dec_device_bufs,
|
||||
host_buffers=self.dec_host_bufs,
|
||||
input_names=self.dec_in_names,
|
||||
output_names=self.dec_out_names,
|
||||
)
|
||||
|
||||
# 3.5) pick best mask
|
||||
best_idx = np.argmax(iou_preds[0]) # shape: (3,)
|
||||
mask_logit = masks[0, best_idx] # shape: (256, 256)
|
||||
mask_prob = 1 / (1 + np.exp(-mask_logit))
|
||||
mask_bin = (mask_prob > 0.5).astype(np.uint8)
|
||||
|
||||
# 4) update internal states
|
||||
self.prev_mask = mask_bin
|
||||
self.prev_feats0 = feats0
|
||||
self.prev_feats1 = feats1
|
||||
self.prev_embed = embed
|
||||
|
||||
print("[track_next_frame] Tracking updated with previous mask as prompt.")
|
||||
return mask_bin
|
||||
@@ -58,7 +58,7 @@ def get_connected_components(mask):
|
||||
- counts: A tensor of shape (N, 1, H, W) containing the area of the connected
|
||||
components for foreground pixels and 0 for background pixels.
|
||||
"""
|
||||
from sam2_realtime import _C
|
||||
from sam2 import _C
|
||||
|
||||
return _C.get_connected_componnets(mask.to(torch.uint8).contiguous())
|
||||
|
||||
|
||||
@@ -4,10 +4,6 @@
|
||||
# This source code is licensed under the license found in the
|
||||
# LICENSE file in the root directory of this source tree.
|
||||
|
||||
import os
|
||||
# Set the CUDA architecture list
|
||||
os.environ["TORCH_CUDA_ARCH_LIST"] = "8.0 8.6+PTX 8.7 9.0 9.0a"
|
||||
|
||||
from setuptools import find_packages, setup
|
||||
from torch.utils.cpp_extension import BuildExtension, CUDAExtension
|
||||
|
||||
@@ -32,18 +28,17 @@ REQUIRED_PACKAGES = [
|
||||
]
|
||||
|
||||
def get_extensions():
|
||||
srcs = ["sam2_realtime/csrc/connected_components.cu"]
|
||||
srcs = ["sam2/csrc/connected_components.cu"]
|
||||
compile_args = {
|
||||
"cxx": [],
|
||||
"nvcc": [
|
||||
"--compiler-bindir=/usr/bin/gcc-12",
|
||||
"-DCUDA_HAS_FP16=1",
|
||||
"-D__CUDA_NO_HALF_OPERATORS__",
|
||||
"-D__CUDA_NO_HALF_CONVERSIONS__",
|
||||
"-D__CUDA_NO_HALF2_OPERATORS__",
|
||||
],
|
||||
}
|
||||
ext_modules = [CUDAExtension("sam2_realtime._C", srcs, extra_compile_args=compile_args)]
|
||||
ext_modules = [CUDAExtension("sam2._C", srcs, extra_compile_args=compile_args)]
|
||||
return ext_modules
|
||||
|
||||
|
||||
@@ -58,7 +53,7 @@ setup(
|
||||
license=LICENSE,
|
||||
packages=find_packages(),
|
||||
install_requires=REQUIRED_PACKAGES,
|
||||
python_requires=">=3.10.15",
|
||||
python_requires=">=3.11.10",
|
||||
ext_modules=get_extensions(),
|
||||
cmdclass={"build_ext": BuildExtension.with_options(no_python_abi_suffix=True)},
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user