Author SHA1 Message Date
Peter Schroedl 233d9a21ce add minimal functionality using tensorrt decoder/encoder engines 2025-01-16 00:48:14 -08:00
PSchroedl 0564c80e07 Merge pull request #7 from pschroedl/revert-6-fix-install 2025-01-13 20:41:31 -08:00
John | Elite Encoder 2a59fe59b3 Revert "Fix node install from conda environments" 2025-01-13 23:10:45 -05:00
PSchroedl 8b00984581 Merge pull request #6 from eliteprox/fix-install 2025-01-09 19:41:11 -08:00
John | Elite Encoder 8db624cfe1 Update README.md 2025-01-09 21:49:26 -05:00
John | Elite Encoder 6258558746 Update README.md 2025-01-09 21:48:56 -05:00
John | Elite Encoder a4c053dae8 Update README.md 2025-01-09 21:34:35 -05:00
John | Elite Encoder 6b2c03f8bf Create pyproject.toml 2025-01-08 23:15:00 -05:00
John | Elite Encoder 0bbc1e0efa Create README.md
add install notes
2025-01-08 23:14:36 -05:00
John | Elite Encoder 60c60c5c2b remove self install from requirements.txt
Resolves issue with conda windows environments losing system environment variable context (e.g. CUDA_HOME error)
2025-01-08 22:56:46 -05:00
PSchroedl 4f587443fb Merge pull request #5 from pschroedl/unbreak_old_workflows_and_speedup
Unbreak old workflows and speedup
2024-12-09 23:58:21 -08:00
Peter Schroedl 37aa0d4c89 add small model option and config 2024-12-10 08:42:30 +01:00
Peter Schroedl 843ca3e733 reduce resolution internally to 512 2024-12-10 08:28:06 +01:00
Peter Schroedl f4e56bd733 make reset_tracking optional 2024-12-10 08:26:06 +01:00
PSchroedl de1fb0ab2a Merge pull request #4 from pschroedl/fix_point_coords
fix: update point coords scaling for mask
2024-12-09 21:47:43 -08:00
7 changed files with 700 additions and 63 deletions
+128 -61
View File
@@ -1,20 +1,25 @@
#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
# Add the directory containing 'sam2_realtime' to sys.path
current_directory = os.path.dirname(os.path.abspath(__file__))
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 comfy.utils import load_torch_file
from sam2_realtime.sam2_tensorrt_predictor import Sam2TensorrtPredictor
from omegaconf import OmegaConf
from hydra.utils import instantiate
@@ -31,7 +36,7 @@ class DownloadAndLoadSAM2RealtimeModel:
def INPUT_TYPES(s):
return {"required": {
"model": ([
'sam2_hiera_tiny.pt',
'sam2_hiera_tiny.pt', 'sam2_hiera_small.pt',
],),
"segmentor": (
['realtime'],
@@ -70,7 +75,8 @@ class DownloadAndLoadSAM2RealtimeModel:
if not os.path.exists(model_path):
print(f"Downloading SAM2 model to: {model_path}")
url = "https://dl.fbaipublicfiles.com/segment_anything_2/072824/sam2_hiera_tiny.pt"
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()
@@ -83,8 +89,9 @@ 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: remove hardcoded config and path
with initialize_config_dir(config_dir=config_dir, version_base=None):
cfg = compose(config_name=model_cfg)
@@ -139,13 +146,13 @@ class Sam2RealtimeSegmentation:
return {
"required": {
"images": ("IMAGE",),
"sam2_model": ("SAM2MODEL",),
"reset_tracking": ("BOOLEAN", {"default": False}),
# "sam2_model": ("SAM2MODEL",),
# "keep_model_loaded": ("BOOLEAN", {"default": True}),
},
"optional": {
"coordinates_positive": ("STRING", ),
"coordinates_negative": ("STRING", ),
# "coordinates_positive": ("STRING", ),
# "coordinates_negative": ("STRING", ),
# "reset_tracking": ("BOOLEAN", {"default": False}),
# "bboxes": ("BBOX", ),
# "individual_objects": ("BOOLEAN", {"default": False}),
# "mask": ("MASK", ),
@@ -158,7 +165,26 @@ class Sam2RealtimeSegmentation:
CATEGORY = "SAM2-Realtime"
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
def _process_coordinate_input(self, coordinates, label):
@@ -190,70 +216,60 @@ class Sam2RealtimeSegmentation:
def segment_images(
self,
images,
sam2_model,
# keep_model_loaded,
reset_tracking,
coordinates_positive=None,
coordinates_negative=None,
#point_labels=None,
# bboxes=None,
# individual_objects=False,
# mask=None,
# sam2_model,
# coordinates_positive=None,
# coordinates_negative=None,
# reset_tracking=False,
):
model = sam2_model["model"]
device = torch.device("cuda")
model.to(device)
processed_frames = []
mask_list = []
# The `model` variable is now ready and equivalent to `predictor` returned by sam2.build_sam.build_sam2_camera_predictor
if reset_tracking:
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)
# if reset_tracking:
# self.if_init = False
with torch.inference_mode(), torch.autocast("cuda", dtype=torch.float16):
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:
self.predictor.load_first_frame(frame)
self.if_init = True
if all_points:
_, _, out_mask_logits = self.predictor.add_new_prompt(
frame_idx=0,
obj_id=1,
points=points_tensor,
labels=labels_tensor,
)
else:
out_mask_logits = torch.zeros((0,), device=device)
out_mask_logits = self.predictor.load_first_frame_and_prompt(
frame_bgr,
point_coord=(512, 512)
)
else:
out_obj_ids, out_mask_logits = self.predictor.track(frame)
out_mask_logits = self.predictor.track(frame_bgr)
# 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
mask_colored = torch.stack([mask] * 3, dim=2)
overlayed_frame = torch.add(frame * 0.7, mask_colored * 0.3)
# # Create colored overlay for processed frames
# mask_colored = torch.stack([mask] * 3, dim=2)
# overlayed_frame = torch.add(frame * 0.7, mask_colored * 0.3)
processed_frames.append(overlayed_frame)
mask_list.append(mask)
# Create colored overlay for processed frames
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
stacked_masks = torch.stack(mask_list, dim=0)
@@ -261,12 +277,63 @@ class Sam2RealtimeSegmentation:
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
"Sam2RealtimeSegmentation": Sam2RealtimeSegmentation,
"Sam2RealtimeSegmentationTest": Sam2RealtimeSegmentationTest
}
NODE_DISPLAY_NAME_MAPPINGS = {
"DownloadAndLoadSAM2RealtimeModel": "(Down)Load sam2_realtime Model",
"Sam2RealtimeSegmentation": "Sam2RealtimeSegmentation"
"Sam2RealtimeSegmentation": "Sam2RealtimeSegmentation",
"Sam2RealtimeSegmentationTest": "Sam2RealtimeSegmentationTest"
}
+10 -1
View File
@@ -1,7 +1,16 @@
pyyaml
ninja
numpy>=1.24.4
tqdm>=4.66.1
hydra-core>=1.3.2
iopath>=0.1.10
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 .
+116
View File
@@ -0,0 +1,116 @@
# @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: 1024
image_size: 512
# 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
+241
View File
@@ -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
+203
View File
@@ -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
+1
View File
@@ -36,6 +36,7 @@ def get_extensions():
compile_args = {
"cxx": [],
"nvcc": [
"--compiler-bindir=/usr/bin/gcc-12",
"-DCUDA_HAS_FP16=1",
"-D__CUDA_NO_HALF_OPERATORS__",
"-D__CUDA_NO_HALF_CONVERSIONS__",