sam2 -> sam2_realtime to disambiguate package imports

This commit is contained in:
Peter Schroedl
2024-11-27 21:30:37 +01:00
parent 6a27c5fb8e
commit 76c2f4f135
25 changed files with 38 additions and 40 deletions
+1
View File
@@ -0,0 +1 @@
__pycache__/
+13 -5
View File
@@ -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"
}
-9
View File
@@ -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")
+5
View File
@@ -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__(