sam2 -> sam2_realtime to disambiguate package imports
This commit is contained in:
@@ -0,0 +1 @@
|
||||
__pycache__/
|
||||
@@ -5,7 +5,15 @@ import numpy as np
|
||||
import logging
|
||||
import json
|
||||
import ast
|
||||
from .sam2.sam2_camera_predictor import SAM2CameraPredictor
|
||||
|
||||
import sys
|
||||
|
||||
# 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)
|
||||
print("sys.path:", sys.path)
|
||||
from sam2_realtime.sam2_tensor_predictor import SAM2TensorPredictor
|
||||
from comfy.utils import load_torch_file
|
||||
|
||||
from omegaconf import OmegaConf
|
||||
@@ -40,7 +48,7 @@ class DownloadAndLoadSAM2RealtimeModel:
|
||||
RETURN_TYPES = ("SAM2MODEL",)
|
||||
RETURN_NAMES = ("sam2_model",)
|
||||
FUNCTION = "loadmodel"
|
||||
CATEGORY = "SAM2-Realtime"
|
||||
CATEGORY = "sam2_realtime"
|
||||
|
||||
def loadmodel(self, model, segmentor, device, precision):
|
||||
if precision != 'fp32' and device == 'cpu':
|
||||
@@ -80,7 +88,7 @@ class DownloadAndLoadSAM2RealtimeModel:
|
||||
cfg = compose(config_name=model_cfg)
|
||||
|
||||
hydra_overrides = [
|
||||
"++model._target_=sam2.sam2_camera_predictor.SAM2CameraPredictor",
|
||||
"++model._target_=sam2_realtime.sam2_tensor_predictor.SAM2TensorPredictor",
|
||||
]
|
||||
hydra_overrides_extra = [
|
||||
"++model.sam_mask_decoder_extra_args.dynamic_multimask_via_stability=true",
|
||||
@@ -146,7 +154,7 @@ class Sam2RealtimeSegmentation:
|
||||
RETURN_NAMES = ("PROCESSED_IMAGES","MASK",)
|
||||
RETURN_TYPES = ("IMAGE", "IMAGE",)
|
||||
FUNCTION = "segment_images"
|
||||
CATEGORY = "SAM2-Realtime"
|
||||
CATEGORY = "sam2_realtime"
|
||||
|
||||
def __init__(self):
|
||||
self.predictor = None
|
||||
@@ -240,6 +248,6 @@ NODE_CLASS_MAPPINGS = {
|
||||
"Sam2RealtimeSegmentation": Sam2RealtimeSegmentation
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"DownloadAndLoadSAM2RealtimeModel": "(Down)Load SAM2-Realtime Model",
|
||||
"DownloadAndLoadSAM2RealtimeModel": "(Down)Load sam2_realtime Model",
|
||||
"Sam2RealtimeSegmentation": "Sam2RealtimeSegmentation"
|
||||
}
|
||||
|
||||
@@ -1,9 +0,0 @@
|
||||
# Copyright (c) Meta Platforms, Inc. and affiliates.
|
||||
# All rights reserved.
|
||||
|
||||
# This source code is licensed under the license found in the
|
||||
# LICENSE file in the root directory of this source tree.
|
||||
|
||||
# from hydra import initialize_config_module
|
||||
|
||||
# initialize_config_module("sam2_configs", version_base="1.2")
|
||||
@@ -0,0 +1,5 @@
|
||||
# Copyright (c) Meta Platforms, Inc. and affiliates.
|
||||
# All rights reserved.
|
||||
|
||||
# This source code is licensed under the license found in the
|
||||
# LICENSE file in the root directory of this source tree.
|
||||
@@ -11,13 +11,13 @@ import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from sam2.modeling.backbones.utils import (
|
||||
from sam2_realtime.modeling.backbones.utils import (
|
||||
PatchEmbed,
|
||||
window_partition,
|
||||
window_unpartition,
|
||||
)
|
||||
|
||||
from sam2.modeling.sam2_utils import DropPath, MLP
|
||||
from sam2_realtime.modeling.sam2_utils import DropPath, MLP
|
||||
|
||||
|
||||
def do_pool(x: torch.Tensor, pool: nn.Module, norm: nn.Module = None) -> torch.Tensor:
|
||||
@@ -9,9 +9,9 @@ from typing import Optional
|
||||
import torch
|
||||
from torch import nn, Tensor
|
||||
|
||||
from sam2.modeling.sam.transformer import RoPEAttention
|
||||
from sam2_realtime.modeling.sam.transformer import RoPEAttention
|
||||
|
||||
from sam2.modeling.sam2_utils import get_activation_fn, get_clones
|
||||
from sam2_realtime.modeling.sam2_utils import get_activation_fn, get_clones
|
||||
|
||||
|
||||
class MemoryAttentionLayer(nn.Module):
|
||||
@@ -11,7 +11,7 @@ import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from sam2.modeling.sam2_utils import DropPath, get_clones, LayerNorm2d
|
||||
from sam2_realtime.modeling.sam2_utils import DropPath, get_clones, LayerNorm2d
|
||||
|
||||
|
||||
class MaskDownSampler(nn.Module):
|
||||
@@ -9,7 +9,7 @@ from typing import List, Optional, Tuple, Type
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
from sam2.modeling.sam2_utils import LayerNorm2d, MLP
|
||||
from sam2_realtime.modeling.sam2_utils import LayerNorm2d, MLP
|
||||
|
||||
|
||||
class MaskDecoder(nn.Module):
|
||||
@@ -9,9 +9,9 @@ from typing import Optional, Tuple, Type
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
from sam2.modeling.position_encoding import PositionEmbeddingRandom
|
||||
from sam2_realtime.modeling.position_encoding import PositionEmbeddingRandom
|
||||
|
||||
from sam2.modeling.sam2_utils import LayerNorm2d
|
||||
from sam2_realtime.modeling.sam2_utils import LayerNorm2d
|
||||
|
||||
|
||||
class PromptEncoder(nn.Module):
|
||||
@@ -13,10 +13,10 @@ import torch
|
||||
import torch.nn.functional as F
|
||||
from torch import nn, Tensor
|
||||
|
||||
from sam2.modeling.position_encoding import apply_rotary_enc, compute_axial_cis
|
||||
from sam2_realtime.modeling.position_encoding import apply_rotary_enc, compute_axial_cis
|
||||
|
||||
from sam2.modeling.sam2_utils import MLP
|
||||
from sam2.utils.misc import get_sdpa_settings
|
||||
from sam2_realtime.modeling.sam2_utils import MLP
|
||||
from sam2_realtime.utils.misc import get_sdpa_settings
|
||||
|
||||
warnings.simplefilter(action="ignore", category=FutureWarning)
|
||||
OLD_GPU, USE_FLASH_ATTN, MATH_KERNEL_ON = get_sdpa_settings()
|
||||
@@ -10,10 +10,10 @@ import torch.nn.functional as F
|
||||
|
||||
from torch.nn.init import trunc_normal_
|
||||
|
||||
from sam2.modeling.sam.mask_decoder import MaskDecoder
|
||||
from sam2.modeling.sam.prompt_encoder import PromptEncoder
|
||||
from sam2.modeling.sam.transformer import TwoWayTransformer
|
||||
from sam2.modeling.sam2_utils import get_1d_sine_pe, MLP, select_closest_cond_frames
|
||||
from sam2_realtime.modeling.sam.mask_decoder import MaskDecoder
|
||||
from sam2_realtime.modeling.sam.prompt_encoder import PromptEncoder
|
||||
from sam2_realtime.modeling.sam.transformer import TwoWayTransformer
|
||||
from sam2_realtime.modeling.sam2_utils import get_1d_sine_pe, MLP, select_closest_cond_frames
|
||||
|
||||
# a large negative value as a placeholder score for missing objects
|
||||
NO_OBJ_SCORE = -1024.0
|
||||
@@ -7,22 +7,15 @@
|
||||
from collections import OrderedDict
|
||||
|
||||
import torch
|
||||
import numpy as np
|
||||
|
||||
from tqdm import tqdm
|
||||
|
||||
import sys
|
||||
import os
|
||||
|
||||
# To import from local sam2
|
||||
sys.path.append(os.path.dirname(os.path.abspath(__file__)))
|
||||
|
||||
from sam2.modeling.sam2_base import NO_OBJ_SCORE, SAM2Base
|
||||
from sam2.utils.misc import concat_points, fill_holes_in_mask_scores, load_video_frames
|
||||
import numpy as np
|
||||
import cv2
|
||||
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 SAM2CameraPredictor(SAM2Base):
|
||||
class SAM2TensorPredictor(SAM2Base):
|
||||
"""The predictor class to handle user interactions and manage inference states."""
|
||||
|
||||
def __init__(
|
||||
Reference in New Issue
Block a user