"""Minimal SAM 3D Objects inference wrapper. Replaces upstream notebook/inference.py, whose module-level imports pull in kaolin.visualize, SceneVisualizer and plotly — none of which are needed to run the pipeline. Import this module only when sam-3d-objects is on sys.path and a GPU context is available. """ import os from typing import Optional, Union import numpy as np from PIL import Image from omegaconf import OmegaConf from hydra.utils import instantiate import sam3d_objects # noqa: F401 guarded by LIDRA_SKIP_INIT # The attention modules read ATTN_BACKEND/SPARSE_ATTN_BACKEND from the # environment exactly once, at import time. inference_pipeline's # set_attention_backend() flips the env to flash_attn on datacenter GPUs, # so the modules must be imported BEFORE the pipeline to stay on sdpa. import sam3d_objects.model.backbone.tdfy_dit.modules.attention # noqa: F401 import sam3d_objects.model.backbone.tdfy_dit.modules.sparse # noqa: F401 from sam3d_objects.pipeline.inference_pipeline_pointmap import InferencePipelinePointMap class SAM3DInference: def __init__(self, config_file: str, compile: bool = False): config = OmegaConf.load(config_file) config.rendering_engine = "pytorch3d" # disable nvdiffrast config.compile_model = compile config.workspace_dir = os.path.dirname(config_file) self._pipeline: InferencePipelinePointMap = instantiate(config) @staticmethod def merge_mask_to_rgba(image: np.ndarray, mask: np.ndarray) -> np.ndarray: mask = mask.astype(np.uint8) * 255 return np.concatenate([image[..., :3], mask[..., None]], axis=-1) def __call__( self, image: Union[Image.Image, np.ndarray], mask: Optional[Union[Image.Image, np.ndarray]], seed: Optional[int] = None, pointmap=None, ) -> dict: image = self.merge_mask_to_rgba(np.asarray(image), np.asarray(mask)) return self._pipeline.run( image, None, seed, stage1_only=False, with_mesh_postprocess=False, with_texture_baking=False, with_layout_postprocess=False, use_vertex_color=True, stage1_inference_steps=None, pointmap=pointmap, )