Compare commits
11
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
233d9a21ce | ||
|
|
0564c80e07 | ||
|
|
2a59fe59b3 | ||
|
|
8b00984581 | ||
|
|
8db624cfe1 | ||
|
|
6258558746 | ||
|
|
a4c053dae8 | ||
|
|
6b2c03f8bf | ||
|
|
0bbc1e0efa | ||
|
|
60c60c5c2b | ||
|
|
4f587443fb |
@@ -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
|
||||
@@ -141,13 +146,13 @@ 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}),
|
||||
# "coordinates_positive": ("STRING", ),
|
||||
# "coordinates_negative": ("STRING", ),
|
||||
# "reset_tracking": ("BOOLEAN", {"default": False}),
|
||||
# "bboxes": ("BBOX", ),
|
||||
# "individual_objects": ("BOOLEAN", {"default": False}),
|
||||
# "mask": ("MASK", ),
|
||||
@@ -160,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):
|
||||
@@ -192,71 +216,60 @@ class Sam2RealtimeSegmentation:
|
||||
def segment_images(
|
||||
self,
|
||||
images,
|
||||
sam2_model,
|
||||
# keep_model_loaded,
|
||||
coordinates_positive=None,
|
||||
coordinates_negative=None,
|
||||
reset_tracking=False,
|
||||
#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)
|
||||
# # 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)
|
||||
# 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)
|
||||
@@ -264,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
@@ -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 .
|
||||
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user