Spaces:
Running on Zero
Running on Zero
Claude Sonnet 4.5
refactor: minimal hacks — own inference wrapper, trimmed stubs and deps, gradio 6.20
7444d60 unverified | """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) | |
| 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, | |
| ) | |