Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
233d9a21ce |
@@ -1,20 +1,25 @@
|
|||||||
|
#nodes.py
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import os
|
import os
|
||||||
import requests
|
import requests
|
||||||
import numpy as np
|
|
||||||
import logging
|
import logging
|
||||||
import json
|
|
||||||
import ast
|
import ast
|
||||||
|
|
||||||
import sys
|
import sys
|
||||||
|
|
||||||
|
import numpy as np #ugh
|
||||||
|
import cv2 #double ugh
|
||||||
|
|
||||||
|
import pycuda.driver as cuda
|
||||||
|
import pycuda.autoinit # Automatically initializes a CUDA context
|
||||||
|
|
||||||
# Add the directory containing 'sam2_realtime' to sys.path
|
# Add the directory containing 'sam2_realtime' to sys.path
|
||||||
current_directory = os.path.dirname(os.path.abspath(__file__))
|
current_directory = os.path.dirname(os.path.abspath(__file__))
|
||||||
sam2_realtime_path = os.path.join(current_directory) # Adjust the relative path
|
sam2_realtime_path = os.path.join(current_directory) # Adjust the relative path
|
||||||
sys.path.append(sam2_realtime_path)
|
sys.path.append(sam2_realtime_path)
|
||||||
|
|
||||||
from sam2_realtime.sam2_tensor_predictor import SAM2TensorPredictor
|
from sam2_realtime.sam2_tensor_predictor import SAM2TensorPredictor
|
||||||
from comfy.utils import load_torch_file
|
from sam2_realtime.sam2_tensorrt_predictor import Sam2TensorrtPredictor
|
||||||
|
|
||||||
from omegaconf import OmegaConf
|
from omegaconf import OmegaConf
|
||||||
from hydra.utils import instantiate
|
from hydra.utils import instantiate
|
||||||
@@ -141,13 +146,13 @@ class Sam2RealtimeSegmentation:
|
|||||||
return {
|
return {
|
||||||
"required": {
|
"required": {
|
||||||
"images": ("IMAGE",),
|
"images": ("IMAGE",),
|
||||||
"sam2_model": ("SAM2MODEL",),
|
# "sam2_model": ("SAM2MODEL",),
|
||||||
# "keep_model_loaded": ("BOOLEAN", {"default": True}),
|
# "keep_model_loaded": ("BOOLEAN", {"default": True}),
|
||||||
},
|
},
|
||||||
"optional": {
|
"optional": {
|
||||||
"coordinates_positive": ("STRING", ),
|
# "coordinates_positive": ("STRING", ),
|
||||||
"coordinates_negative": ("STRING", ),
|
# "coordinates_negative": ("STRING", ),
|
||||||
"reset_tracking": ("BOOLEAN", {"default": False}),
|
# "reset_tracking": ("BOOLEAN", {"default": False}),
|
||||||
# "bboxes": ("BBOX", ),
|
# "bboxes": ("BBOX", ),
|
||||||
# "individual_objects": ("BOOLEAN", {"default": False}),
|
# "individual_objects": ("BOOLEAN", {"default": False}),
|
||||||
# "mask": ("MASK", ),
|
# "mask": ("MASK", ),
|
||||||
@@ -160,7 +165,26 @@ class Sam2RealtimeSegmentation:
|
|||||||
CATEGORY = "SAM2-Realtime"
|
CATEGORY = "SAM2-Realtime"
|
||||||
|
|
||||||
def __init__(self):
|
def __init__(self):
|
||||||
self.predictor = None
|
##############################################################################
|
||||||
|
# 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.if_init = False
|
self.if_init = False
|
||||||
|
|
||||||
def _process_coordinate_input(self, coordinates, label):
|
def _process_coordinate_input(self, coordinates, label):
|
||||||
@@ -192,71 +216,60 @@ class Sam2RealtimeSegmentation:
|
|||||||
def segment_images(
|
def segment_images(
|
||||||
self,
|
self,
|
||||||
images,
|
images,
|
||||||
sam2_model,
|
# sam2_model,
|
||||||
# keep_model_loaded,
|
# coordinates_positive=None,
|
||||||
coordinates_positive=None,
|
# coordinates_negative=None,
|
||||||
coordinates_negative=None,
|
# reset_tracking=False,
|
||||||
reset_tracking=False,
|
|
||||||
#point_labels=None,
|
|
||||||
# bboxes=None,
|
|
||||||
# individual_objects=False,
|
|
||||||
# mask=None,
|
|
||||||
):
|
):
|
||||||
model = sam2_model["model"]
|
|
||||||
device = torch.device("cuda")
|
|
||||||
model.to(device)
|
|
||||||
|
|
||||||
processed_frames = []
|
processed_frames = []
|
||||||
mask_list = []
|
mask_list = []
|
||||||
# The `model` variable is now ready and equivalent to `predictor` returned by sam2.build_sam.build_sam2_camera_predictor
|
|
||||||
|
|
||||||
if reset_tracking:
|
# if reset_tracking:
|
||||||
self.if_init = False
|
# self.if_init = False
|
||||||
self.predictor = None
|
|
||||||
|
|
||||||
if self.predictor is None:
|
|
||||||
self.predictor = model
|
|
||||||
|
|
||||||
# Process coordinates once, outside the frame loop
|
|
||||||
pos_points, pos_labels = self._process_coordinate_input(coordinates_positive, 1)
|
|
||||||
neg_points, neg_labels = self._process_coordinate_input(coordinates_negative, 0)
|
|
||||||
all_points = pos_points + neg_points
|
|
||||||
all_labels = pos_labels + neg_labels
|
|
||||||
|
|
||||||
if all_points:
|
|
||||||
points_tensor = torch.tensor([all_points], device=device)
|
|
||||||
labels_tensor = torch.tensor([all_labels], device=device)
|
|
||||||
|
|
||||||
with torch.inference_mode(), torch.autocast("cuda", dtype=torch.float16):
|
with torch.inference_mode(), torch.autocast("cuda", dtype=torch.float16):
|
||||||
for frame_idx, frame in enumerate(images):
|
for frame_idx, frame in enumerate(images):
|
||||||
frame = frame.to(device).float()
|
|
||||||
|
# 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)
|
||||||
|
|
||||||
if not self.if_init:
|
if not self.if_init:
|
||||||
self.predictor.load_first_frame(frame)
|
|
||||||
self.if_init = True
|
self.if_init = True
|
||||||
|
|
||||||
if all_points:
|
out_mask_logits = self.predictor.load_first_frame_and_prompt(
|
||||||
_, _, out_mask_logits = self.predictor.add_new_prompt(
|
frame_bgr,
|
||||||
frame_idx=0,
|
point_coord=(512, 512)
|
||||||
obj_id=1,
|
)
|
||||||
points=points_tensor,
|
|
||||||
labels=labels_tensor,
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
out_mask_logits = torch.zeros((0,), device=device)
|
|
||||||
else:
|
else:
|
||||||
out_obj_ids, out_mask_logits = self.predictor.track(frame)
|
out_mask_logits = self.predictor.track(frame_bgr)
|
||||||
|
|
||||||
# Process mask logits
|
# Process mask logits
|
||||||
mask = self._process_mask_logits(out_mask_logits, frame.shape, device)
|
# mask = self._process_mask_logits(out_mask_logits, frame.shape, self.device)
|
||||||
|
mask = out_mask_logits
|
||||||
|
|
||||||
# Create colored overlay for processed frames
|
# # Create colored overlay for processed frames
|
||||||
mask_colored = torch.stack([mask] * 3, dim=2)
|
# mask_colored = torch.stack([mask] * 3, dim=2)
|
||||||
|
|
||||||
overlayed_frame = torch.add(frame * 0.7, mask_colored * 0.3)
|
# overlayed_frame = torch.add(frame * 0.7, mask_colored * 0.3)
|
||||||
|
|
||||||
processed_frames.append(overlayed_frame)
|
# Create colored overlay for processed frames
|
||||||
mask_list.append(mask)
|
mask_colored = np.stack([mask] * 3, axis=2) # Stack along the last dimension (HxWxC format)
|
||||||
|
|
||||||
|
mask_colored_resized = cv2.resize(mask_colored, (frame.shape[1], frame.shape[0]))
|
||||||
|
|
||||||
|
# Overlay the frame and the colored mask
|
||||||
|
overlaid_frame = frame * 0.7 + mask_colored_resized * 0.3
|
||||||
|
|
||||||
|
processed_frames.append(torch.tensor(overlaid_frame))
|
||||||
|
mask_list.append(torch.tensor(mask))
|
||||||
|
|
||||||
# Stack masks and frames
|
# Stack masks and frames
|
||||||
stacked_masks = torch.stack(mask_list, dim=0)
|
stacked_masks = torch.stack(mask_list, dim=0)
|
||||||
@@ -264,12 +277,63 @@ class Sam2RealtimeSegmentation:
|
|||||||
|
|
||||||
return (stacked_frames, stacked_masks)
|
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 = {
|
NODE_CLASS_MAPPINGS = {
|
||||||
"DownloadAndLoadSAM2RealtimeModel": DownloadAndLoadSAM2RealtimeModel,
|
"DownloadAndLoadSAM2RealtimeModel": DownloadAndLoadSAM2RealtimeModel,
|
||||||
"Sam2RealtimeSegmentation": Sam2RealtimeSegmentation
|
"Sam2RealtimeSegmentation": Sam2RealtimeSegmentation,
|
||||||
|
"Sam2RealtimeSegmentationTest": Sam2RealtimeSegmentationTest
|
||||||
}
|
}
|
||||||
|
|
||||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||||
"DownloadAndLoadSAM2RealtimeModel": "(Down)Load sam2_realtime Model",
|
"DownloadAndLoadSAM2RealtimeModel": "(Down)Load sam2_realtime Model",
|
||||||
"Sam2RealtimeSegmentation": "Sam2RealtimeSegmentation"
|
"Sam2RealtimeSegmentation": "Sam2RealtimeSegmentation",
|
||||||
|
"Sam2RealtimeSegmentationTest": "Sam2RealtimeSegmentationTest"
|
||||||
}
|
}
|
||||||
|
|||||||
+10
-1
@@ -1,7 +1,16 @@
|
|||||||
pyyaml
|
pyyaml
|
||||||
|
ninja
|
||||||
numpy>=1.24.4
|
numpy>=1.24.4
|
||||||
tqdm>=4.66.1
|
tqdm>=4.66.1
|
||||||
hydra-core>=1.3.2
|
hydra-core>=1.3.2
|
||||||
iopath>=0.1.10
|
iopath>=0.1.10
|
||||||
pillow>=9.4.0
|
pillow>=9.4.0
|
||||||
git+https://github.com/pschroedl/ComfyUI-SAM2-Realtime.git@main#egg=sam2_realtime
|
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 .
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,241 @@
|
|||||||
|
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
|
||||||
@@ -0,0 +1,203 @@
|
|||||||
|
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
|
||||||
@@ -36,6 +36,7 @@ def get_extensions():
|
|||||||
compile_args = {
|
compile_args = {
|
||||||
"cxx": [],
|
"cxx": [],
|
||||||
"nvcc": [
|
"nvcc": [
|
||||||
|
"--compiler-bindir=/usr/bin/gcc-12",
|
||||||
"-DCUDA_HAS_FP16=1",
|
"-DCUDA_HAS_FP16=1",
|
||||||
"-D__CUDA_NO_HALF_OPERATORS__",
|
"-D__CUDA_NO_HALF_OPERATORS__",
|
||||||
"-D__CUDA_NO_HALF_CONVERSIONS__",
|
"-D__CUDA_NO_HALF_CONVERSIONS__",
|
||||||
|
|||||||
Reference in New Issue
Block a user