import spaces import os import sys import pickle import random import shutil # Path setup MUST happen before any projects.sa2va... import SPACE_ROOT = os.path.dirname(os.path.abspath(__file__)) for _p in (SPACE_ROOT,): if _p not in sys.path: sys.path.insert(0, _p) import numpy as np import torch import torch.nn.functional as F import torchvision.transforms as transforms from PIL import Image from torchvision.transforms.functional import to_pil_image from transformers import ( AutoTokenizer, Qwen2_5_VLProcessor, ) from peft import LoraConfig from xtuner.utils import PROMPT_TEMPLATE, IGNORE_INDEX from xtuner.registry import BUILDER from projects.sa2va.models.sa2va import Sa2VAModel from projects.sa2va.models.sam2_train import SAM2TrainRunner from projects.sa2va.models.mllm.qwenvl import Qwen2_5_VL from projects.sa2va.datasets.data_utils import sa2va_collect_fn_multitask from projects.sa2va.datasets.base import Sa2VABaseDataset from projects.sa2va.datasets.common import ANSWER_LIST from utils.hm_utils import add_star_marker from third_parts.mmdet.models.losses.cross_entropy_loss import CrossEntropyLoss from third_parts.mmdet.models.losses.dice_loss import DiceLoss # --------------------------------------------------------------------------- # Prompt templates # --------------------------------------------------------------------------- TASK_PROMPT = ( "Regions with same base material but different colors are considered as " "different materials. However, regions with different lighting, shading " "or shadows are considered as the same material." ) STAR_QUESTIONS = [ f"Please segment all pixels with the same material as where the star is. {TASK_PROMPT}", f"Can you segment all pixels with the same material where the star is located? {TASK_PROMPT}", f"Segment all areas that have the same material as where the star is. {TASK_PROMPT}", ] REFERRING_QUESTIONS = [ f"Please segment all pixels made of the material described below. {TASK_PROMPT}\nDescription: ", f"Can you segment all pixels that match the material described below? {TASK_PROMPT}\nDescription: ", f"Segment every region that has the material described below. {TASK_PROMPT}\nDescription: ", ] SEG_QUESTIONS = [ "Can you segment the {class_name} in this image?", "Please segment {class_name} in this image.", "Could you provide a segmentation mask for the {class_name} in this image?", ] # --------------------------------------------------------------------------- # Constants # --------------------------------------------------------------------------- QWEN_MODEL_PATH = "Qwen/Qwen2.5-VL-7B-Instruct" MAOAM_CKPT_REPO = "jpark677/maoam_ckpts" MAOAM_CKPT_FILE = "sa2va/mp_rank_00_model_states.pt" SAM2_CKPT_REPO = "facebook/sam2-hiera-large" SAM2_CKPT_FILE = "sam2_hiera_large.pt" QWEN_MIN_PIXELS = 512 * 28 * 28 QWEN_MAX_PIXELS = 2048 * 28 * 28 # --------------------------------------------------------------------------- # Download SAM2 checkpoint # --------------------------------------------------------------------------- sam2_dir = os.path.join(SPACE_ROOT, "pretrained", "sam2") os.makedirs(sam2_dir, exist_ok=True) sam2_ckpt_local = os.path.join(sam2_dir, SAM2_CKPT_FILE) if not os.path.exists(sam2_ckpt_local): from huggingface_hub import hf_hub_download print(f"[INFO] Downloading SAM2 checkpoint from {SAM2_CKPT_REPO}...") sam2_ckpt_path = hf_hub_download(SAM2_CKPT_REPO, SAM2_CKPT_FILE, repo_type="model") shutil.copy(sam2_ckpt_path, sam2_ckpt_local) print(f"[INFO] SAM2 checkpoint saved to {sam2_ckpt_local}") # --------------------------------------------------------------------------- # Safe checkpoint loading # --------------------------------------------------------------------------- def _torch_load_ckpt_safely(path: str): try: return torch.load(path, map_location="cpu", weights_only=True) except (pickle.UnpicklingError, TypeError): print(f"[WARN] weights_only=True failed, retrying with weights_only=False") try: return torch.load(path, map_location="cpu", weights_only=False) except TypeError: return torch.load(path, map_location="cpu") def _load_state_dict_from_mp_rank(path: str) -> dict: ckpt = _torch_load_ckpt_safely(path) if isinstance(ckpt, dict): if "module" in ckpt and isinstance(ckpt["module"], dict): state_dict = ckpt["module"] elif "state_dict" in ckpt and isinstance(ckpt["state_dict"], dict): state_dict = ckpt["state_dict"] else: state_dict = ckpt else: state_dict = ckpt if not isinstance(state_dict, dict): raise ValueError(f"Unsupported checkpoint format at {path}") return state_dict # --------------------------------------------------------------------------- # Model building # --------------------------------------------------------------------------- def build_model(): """Build the Sa2VAModel with MAOAM weights.""" special_tokens = ["[SEG]", "

", "

", "", ""] tokenizer_cfg = dict( type=AutoTokenizer.from_pretrained, pretrained_model_name_or_path=QWEN_MODEL_PATH, trust_remote_code=True, padding_side="right", ) model_cfg = dict( type=Sa2VAModel, training_bs=1, special_tokens=special_tokens, pretrained_pth=None, fix_number=1, loss_sample_points=True, frozen_sam2_decoder=False, arch_type="qwen", weight_star=1.0, weight_referring=1.0, weight_vqa=0.0, mllm=dict( type=Qwen2_5_VL, model_path=QWEN_MODEL_PATH, freeze_llm=True, freeze_visual_encoder=True, llm_lora=dict( type=LoraConfig, r=128, lora_alpha=256, lora_dropout=0.05, bias="none", task_type="CAUSAL_LM", modules_to_save=["lm_head", "embed_tokens"], target_modules=None, ), ), tokenizer=tokenizer_cfg, grounding_encoder=dict( type=SAM2TrainRunner, ckpt_path=os.path.abspath(sam2_ckpt_local), ), loss_mask=dict( type=CrossEntropyLoss, use_sigmoid=True, reduction="mean", loss_weight=2.0, ), loss_dice=dict( type=DiceLoss, use_sigmoid=True, activate=True, reduction="mean", naive_dice=True, eps=1.0, loss_weight=0.5, ), ) print("[INFO] Building Sa2VAModel...") sa2va_model = BUILDER.build(model_cfg) # Download and load MAOAM checkpoint from huggingface_hub import hf_hub_download print(f"[INFO] Downloading MAOAM checkpoint from {MAOAM_CKPT_REPO}...") maoam_ckpt_path = hf_hub_download( MAOAM_CKPT_REPO, MAOAM_CKPT_FILE, repo_type="model" ) state_dict = _load_state_dict_from_mp_rank(maoam_ckpt_path) # Strip common DDP prefix if len(state_dict) > 0 and all(k.startswith("module.") for k in state_dict.keys()): state_dict = {k[len("module."):]: v for k, v in state_dict.items()} missing, unexpected = sa2va_model.load_state_dict(state_dict, strict=False) print(f"[INFO] Loaded MAOAM checkpoint: missing={len(missing)} unexpected={len(unexpected)}") # Force update lm_head weight (critical for Qwen untied embeddings) lm_head_key = "mllm.model.lm_head.weight" if lm_head_key in state_dict: lm_head_weight = state_dict[lm_head_key] sa2va_model.mllm.model.get_output_embeddings().weight.data.copy_(lm_head_weight) print("[INFO] Force updated lm_head weight from pretrained state_dict.") sa2va_model = sa2va_model.eval() sa2va_model.to("cuda") print("[INFO] Model loaded and moved to CUDA.") return sa2va_model sa2va_model = build_model() # --------------------------------------------------------------------------- # Build interactive packer # --------------------------------------------------------------------------- tokenizer_cfg = dict( type=AutoTokenizer.from_pretrained, pretrained_model_name_or_path=QWEN_MODEL_PATH, trust_remote_code=True, padding_side="right", ) preprocessor_cfg = dict( type=Qwen2_5_VLProcessor.from_pretrained, pretrained_model_name_or_path=QWEN_MODEL_PATH, trust_remote_code=True, ) class _InteractivePacker(Sa2VABaseDataset): def real_len(self): return 1 def prepare_data(self, index): raise NotImplementedError def process_qwen_image(self, img_chw_float, min_pixels, max_pixels): img_chw_float = img_chw_float.clamp(0.0, 1.0) pil = to_pil_image(img_chw_float) merge_length = self.preprocessor.image_processor.merge_size ** 2 d = self.preprocessor.image_processor( images=[pil], min_pixels=int(min_pixels), max_pixels=int(max_pixels), ) pixel_values = torch.as_tensor(d["pixel_values"], dtype=torch.float32) image_grid_thw = torch.as_tensor(d["image_grid_thw"], dtype=torch.long) num_image_tokens = int(image_grid_thw[0].prod().item()) // int(merge_length) return pixel_values, image_grid_thw, num_image_tokens def pack_task_qwen(self, qwen_image_chw_float, question, answer, qwen_min_pixels, qwen_max_pixels): pixel_values, image_grid_thw, num_image_tokens = self.process_qwen_image( qwen_image_chw_float, min_pixels=qwen_min_pixels, max_pixels=qwen_max_pixels ) image_token_str = self._create_image_token_string(num_image_tokens) conv = [ {"from": "human", "value": question}, {"from": "gpt", "value": answer}, ] conv = self._process_conversations_for_encoding(conv, image_token_str=image_token_str, is_video=False) conv_prompt = conv[0]["input"] if len(conv) > 0 and "input" in conv[0] else "" token_dict = self.get_inputid_labels(conv) return { "input_ids": token_dict["input_ids"], "labels": token_dict["labels"], "pixel_values": pixel_values, "image_grid_thw": image_grid_thw, "convs": conv_prompt, "question": question, } interactive_packer = _InteractivePacker( tokenizer=tokenizer_cfg, prompt_template=PROMPT_TEMPLATE.qwen_chat, max_length=8192, special_tokens=["[SEG]", "

", "

", "", ""], arch_type="qwen", preprocessor=preprocessor_cfg, repeats=1.0, name="InteractivePacker", ) interactive_packer.tokenizer.add_tokens(["[SEG]", "

", "

", "", ""], special_tokens=True) print("[INFO] Interactive packer ready.") # --------------------------------------------------------------------------- # Image helpers # --------------------------------------------------------------------------- def _ensure_unit_range(tensor): if tensor.numel() == 0: return tensor tensor = tensor.to(dtype=torch.float32) t_min, t_max = float(tensor.min().item()), float(tensor.max().item()) if 0.0 <= t_min and t_max <= 1.0: return tensor if 0.0 <= t_min and t_max <= 255.0: tensor = tensor / 255.0 else: denom = max(t_max - t_min, 1e-6) tensor = (tensor - t_min) / denom return tensor.clamp_(0.0, 1.0) def _image_to_tensor(image): if isinstance(image, torch.Tensor): image_tensor = image.detach().clone() if image_tensor.ndim == 3 and image_tensor.shape[0] in (1, 3): pass elif image_tensor.ndim == 3: image_tensor = image_tensor.permute(2, 0, 1) elif image_tensor.ndim == 2: image_tensor = image_tensor.unsqueeze(0) image_tensor = image_tensor.to(dtype=torch.float32) elif isinstance(image, np.ndarray): if image.dtype == np.uint8: pil_image = Image.fromarray(image) image_tensor = transforms.ToTensor()(pil_image) else: image_tensor = torch.from_numpy(image).float() if image_tensor.ndim == 3 and image_tensor.shape[2] in (1, 3): image_tensor = image_tensor.permute(2, 0, 1) elif image_tensor.ndim == 2: image_tensor = image_tensor.unsqueeze(0) image_tensor = image_tensor / 255.0 if image_tensor.max() > 1.0 else image_tensor else: image_tensor = transforms.ToTensor()(image) return _ensure_unit_range(image_tensor) def resize_image_to_square(image, target_size=1024): image_tensor = _image_to_tensor(image) h, w = image_tensor.shape[-2:] if h < w: new_h = target_size new_w = int(target_size * w / h) else: new_w = target_size new_h = int(target_size * h / w) image_tensor = F.interpolate( image_tensor.unsqueeze(0), size=(new_h, new_w), mode="bilinear", align_corners=False, ).squeeze(0) image_tensor = transforms.CenterCrop((target_size, target_size))(image_tensor) return image_tensor def _overlay_all_stars_1024(base_tensor_1024_chw, coords_1024, fixed_color): img = base_tensor_1024_chw.clone() latched = fixed_color marker_size = max(8, int(1024 // 32)) for i, (hh, ww) in enumerate(coords_1024): h = int(max(0, min(1023, hh))) w = int(max(0, min(1023, ww))) try: if i == 0 and latched is None: img, c = add_star_marker(img, h, w, size=marker_size) latched = c or "blue" else: img, _ = add_star_marker(img, h, w, size=marker_size, color=latched) except TypeError: img, c = add_star_marker(img, h, w, size=marker_size) if i == 0 and latched is None: latched = c or "blue" return img, latched def _create_cyan_overlay(base_rgb_uint8, mask_bool, alpha=0.45): base = base_rgb_uint8.astype(np.float32) overlay = base.copy() cyan = np.array([0.0, 255.0, 255.0], dtype=np.float32) overlay[mask_bool] = overlay[mask_bool] * (1.0 - alpha) + cyan * alpha return np.clip(overlay, 0, 255).astype(np.uint8) # --------------------------------------------------------------------------- # Inference # --------------------------------------------------------------------------- @spaces.GPU(duration=120) def run_inference(clean_pil_1024, coords_1024, text_prompt, fixed_color, disp_tensor_chw): if clean_pil_1024 is None: return None, None, "Please upload an image first." coords_1024 = coords_1024 or [] used_color = fixed_color if len(coords_1024) == 0: final_prompt = text_prompt.replace("", "").replace("", "the material").replace(" ", " ").strip() model_image = clean_pil_1024 task_key = "referring" else: if disp_tensor_chw is None: base_chw = _image_to_tensor(clean_pil_1024) else: base_chw = disp_tensor_chw if used_color is None: _, used_color = _overlay_all_stars_1024(_image_to_tensor(clean_pil_1024), coords_1024, None) final_prompt = text_prompt.replace("", used_color or "blue").replace("", "the material") model_image = transforms.ToPILImage()(base_chw) task_key = "star" clean_chw = _image_to_tensor(clean_pil_1024) g_u8 = (clean_chw.clamp(0.0, 1.0) * 255.0).to(torch.uint8) dummy_mask = torch.zeros((1, 1024, 1024), dtype=torch.uint8) tasks = {} question_for_model = "\n" + final_prompt.strip() answer = random.choice(ANSWER_LIST) if len(coords_1024) == 0: tasks["referring"] = interactive_packer.pack_task_qwen( qwen_image_chw_float=clean_chw, question=question_for_model, answer=answer, qwen_min_pixels=QWEN_MIN_PIXELS, qwen_max_pixels=QWEN_MAX_PIXELS, ) instance = { "src": "gradio", "images_without_star": clean_chw, "g_pixel_values": g_u8, "masks": dummy_mask, "tasks": tasks, } else: star_chw = _image_to_tensor(model_image) tasks["star"] = interactive_packer.pack_task_qwen( qwen_image_chw_float=star_chw, question=question_for_model, answer=answer, qwen_min_pixels=QWEN_MIN_PIXELS, qwen_max_pixels=QWEN_MAX_PIXELS, ) instance = { "src": "gradio", "images_star": star_chw, "g_pixel_values": g_u8, "masks": dummy_mask, "tasks": tasks, } batch = sa2va_collect_fn_multitask([instance])["data"] batch["inference"] = True with torch.no_grad(): out = sa2va_model(batch, None, mode="loss") task_out = out.get(task_key, {}) pred_masks = task_out.get("pred_masks", []) pred_masks_np = [] for m in pred_masks: if torch.is_tensor(m): pred_masks_np.append(m.detach().cpu().numpy()) else: pred_masks_np.append(np.asarray(m)) orig_np = np.array(clean_pil_1024.convert("RGB")) if clean_pil_1024 else None if not pred_masks_np or orig_np is None: return orig_np, None, f"No mask produced. Task: {task_key}" m = pred_masks_np[0] if m.ndim == 3: m = m[0] mb = (m > 0.5) if m.dtype != np.uint8 else (m > 0) if mb.shape != orig_np.shape[:2]: mb_resized = np.array( Image.fromarray((mb * 255).astype(np.uint8)).resize( (orig_np.shape[1], orig_np.shape[0]), Image.NEAREST ) ) > 0 mb = mb_resized overlay_np = _create_cyan_overlay(orig_np, mb) binary_np = (mb.astype(np.uint8) * 255) status = f"Task: {task_key} | Color: {used_color or 'N/A'} | Prompt: {final_prompt}" return overlay_np, binary_np, status # --------------------------------------------------------------------------- # Gradio UI # --------------------------------------------------------------------------- import gradio as gr MAX_STARS = 5 _SELECTION_MODES = { "Material: click": ( STAR_QUESTIONS[0], "Place one or more stars on the image, then click **Submit**. `` is filled automatically.", ), "Material: text": ( REFERRING_QUESTIONS[0], "Replace **``** with your material description (e.g. *shiny chrome metal*). No stars needed.", ), "Object: text": ( SEG_QUESTIONS[0].replace("{class_name}", ""), "Replace **``** with your object expression (e.g. *the man in a red shirt*). No stars needed.", ), } _DEFAULT_MODE = "Material: click" def on_image_upload(image_np): if image_np is None: return None, [], None, None, None, None, "Upload an image to start." orig_pil = Image.fromarray(image_np.astype(np.uint8)).convert("RGB") clean_chw = resize_image_to_square(orig_pil, 1024) disp_np = np.array(transforms.ToPILImage()(clean_chw)) clean_pil_1024 = transforms.ToPILImage()(clean_chw) return disp_np, [], None, clean_pil_1024, clean_chw, clean_chw, "Image loaded. Click up to 5 points, then Submit." def on_click_add_star(image_disp_np, coords_1024, fixed_color, disp_tensor_chw, evt: gr.SelectData): if image_disp_np is None or disp_tensor_chw is None: return image_disp_np, coords_1024, fixed_color, disp_tensor_chw, "Please upload an image first." if evt is None: return image_disp_np, coords_1024, fixed_color, disp_tensor_chw, "Click anywhere on the image to add a star." if coords_1024 is None: coords_1024 = [] if len(coords_1024) >= MAX_STARS: return image_disp_np, coords_1024, fixed_color, disp_tensor_chw, f"Max {MAX_STARS} stars reached." h, w = int(evt.index[1]), int(evt.index[0]) disp = disp_tensor_chw.clone() marker_size = 32 try: if fixed_color is None: disp, c = add_star_marker(disp, h, w, size=marker_size) fixed_color = c or "blue" else: disp, _ = add_star_marker(disp, h, w, size=marker_size, color=fixed_color) except TypeError: disp, c = add_star_marker(disp, h, w, size=marker_size) if fixed_color is None: fixed_color = c or "blue" coords_1024 = coords_1024 + [(h, w)] disp_np = np.array(transforms.ToPILImage()(disp)) return disp_np, coords_1024, fixed_color, disp, f"Star #{len(coords_1024)} @ (h={h}, w={w})." def on_undo_last(coords_1024, fixed_color, clean_tensor_chw): if clean_tensor_chw is None: return None, coords_1024, fixed_color, None, "Nothing to undo." if not coords_1024: disp = clean_tensor_chw.clone() return np.array(transforms.ToPILImage()(disp)), [], None, disp, "Nothing to undo." new_coords = coords_1024[:-1] disp = clean_tensor_chw.clone() latched = fixed_color marker_size = max(8, int(1024 // 32)) for i, (h, w) in enumerate(new_coords): try: if i == 0 and latched is None: disp, c = add_star_marker(disp, int(h), int(w), size=marker_size) latched = c or "blue" else: disp, _ = add_star_marker(disp, int(h), int(w), size=marker_size, color=latched) except TypeError: disp, c = add_star_marker(disp, int(h), int(w), size=marker_size) if i == 0 and latched is None: latched = c or "blue" return np.array(transforms.ToPILImage()(disp)), new_coords, latched, disp, f"Removed last star. {len(new_coords)} remaining." def on_clear_stars(clean_tensor_chw): if clean_tensor_chw is None: return None, [], None, None, "Nothing to clear." disp = clean_tensor_chw.clone() return np.array(transforms.ToPILImage()(disp)), [], None, disp, "Cleared all stars." def on_submit(orig_pil_1024, coords_1024, text_prompt, fixed_color, disp_tensor_chw): if orig_pil_1024 is None: return None, None, "Please upload an image first." return run_inference(orig_pil_1024, coords_1024, text_prompt, fixed_color, disp_tensor_chw) def on_selection_change(mode): prompt, hint = _SELECTION_MODES.get(mode, _SELECTION_MODES[_DEFAULT_MODE]) return prompt, hint with gr.Blocks(theme=gr.themes.Citrus(), title="MAOAM Demo") as demo: gr.Markdown("# MAOAM: Unified Object and Material Selection") gr.Markdown( "Upload an image, choose a selection mode, and get a pixel-accurate segmentation mask. " "Click on the image to place star markers for material selection, or use text prompts." ) coords_state = gr.State([]) fixed_color_state = gr.State(None) orig_pil_state = gr.State(None) # clean PIL 1024-square (for inference) clean_tensor_state = gr.State(None) # CHW clean 1024 torch.Tensor [0,1] disp_tensor_state = gr.State(None) # CHW display tensor with stars with gr.Row(): with gr.Column(scale=1): input_image = gr.Image(label="Input / Click to add star(s)", type="numpy", height=400) selection_dropdown = gr.Dropdown( choices=list(_SELECTION_MODES.keys()), value=_DEFAULT_MODE, label="Selection type", ) text_prompt = gr.Textbox( label="Text prompt", value=_SELECTION_MODES[_DEFAULT_MODE][0], lines=3, ) hint_md = gr.Markdown(_SELECTION_MODES[_DEFAULT_MODE][1]) with gr.Row(): undo_btn = gr.Button("Undo last star", variant="secondary") clear_btn = gr.Button("Clear stars", variant="secondary") submit_btn = gr.Button("Submit", variant="primary") with gr.Column(scale=1): overlay_image = gr.Image(label="Overlaid Image", height=400) binary_mask_image = gr.Image(label="Binary Mask", height=400) status_text = gr.Textbox( label="Status", value="Upload an image, click up to 5 star points, then Submit.", interactive=False, ) coords_table = gr.Dataframe( headers=["h", "w"], datatype=["number", "number"], row_count=5, col_count=(2, "fixed"), interactive=False, label="Star coordinates (1024 space)", ) input_image.upload( on_image_upload, inputs=[input_image], outputs=[input_image, coords_state, fixed_color_state, orig_pil_state, clean_tensor_state, disp_tensor_state, status_text], ).then( lambda coords: [[h, w] for (h, w) in (coords or [])], inputs=[coords_state], outputs=[coords_table], ) input_image.select( on_click_add_star, inputs=[input_image, coords_state, fixed_color_state, disp_tensor_state], outputs=[input_image, coords_state, fixed_color_state, disp_tensor_state, status_text], ).then( lambda coords: [[h, w] for (h, w) in (coords or [])], inputs=[coords_state], outputs=[coords_table], ) undo_btn.click( on_undo_last, inputs=[coords_state, fixed_color_state, clean_tensor_state], outputs=[input_image, coords_state, fixed_color_state, disp_tensor_state, status_text], ).then( lambda coords: [[h, w] for (h, w) in (coords or [])], inputs=[coords_state], outputs=[coords_table], ) clear_btn.click( on_clear_stars, inputs=[clean_tensor_state], outputs=[input_image, coords_state, fixed_color_state, disp_tensor_state, status_text], ).then( lambda coords: [[h, w] for (h, w) in (coords or [])], inputs=[coords_state], outputs=[coords_table], ) submit_btn.click( on_submit, inputs=[orig_pil_state, coords_state, text_prompt, fixed_color_state, disp_tensor_state], outputs=[overlay_image, binary_mask_image, status_text], ) selection_dropdown.change( on_selection_change, inputs=[selection_dropdown], outputs=[text_prompt, hint_md], ) def on_example_click(image_np, mode, prompt): if image_np is not None: disp_np, _, _, clean_pil, clean_chw, disp_chw, status = on_image_upload(image_np) return disp_np, [], None, clean_pil, clean_chw, disp_chw, prompt, _SELECTION_MODES.get(mode, _SELECTION_MODES[_DEFAULT_MODE])[1] return image_np, [], None, None, None, None, prompt, _SELECTION_MODES.get(mode, _SELECTION_MODES[_DEFAULT_MODE])[1] gr.Examples( examples=[ [os.path.join(SPACE_ROOT, "examples", "landscape.jpg"), "Material: click", STAR_QUESTIONS[0]], [os.path.join(SPACE_ROOT, "examples", "dog.jpg"), "Object: text", "Can you segment the dog in this image?"], [os.path.join(SPACE_ROOT, "examples", "cake.jpg"), "Material: text", REFERRING_QUESTIONS[0]], [os.path.join(SPACE_ROOT, "examples", "tent.jpg"), "Material: click", STAR_QUESTIONS[0]], ], inputs=[input_image, selection_dropdown, text_prompt], outputs=[input_image, coords_state, fixed_color_state, orig_pil_state, clean_tensor_state, disp_tensor_state, text_prompt, hint_md], fn=on_example_click, cache_examples=False, run_on_click=True, ) demo.queue() demo.launch()