import gradio as gr import time import os from DataSets.getdata import thermal_feet_dataset from segment_anything1.config import SAM1_MODELS from segment_anything2.config import SAM2_MODELS from groundino_samnet.model import load_models, GSamnet from groundino_samnet.visuals import show_result import io examples = ["assets/t0.jpg","assets/outputi.png"] def load_model_def(): progress = gr.Progress() progress(0) dino = load_models("dino") progress(50) sam = load_models("sam2_l") progress(100) time.sleep(0.5) return dino, sam def predict(model, image): return model([image]) def update_model(selected_model, dino_model): sam_model = load_models(selected_model) sam_args = {"model": sam_model, "points": True, "torch": True} dino_args = {"model": dino_model} modelf = GSamnet(dino_args=dino_args, sam_args=sam_args).to(sam_model.device) return modelf def process_image(uploaded_image, model, om): if uploaded_image is not None: img_byte_arr = io.BytesIO() uploaded_image.save(img_byte_arr, format='PNG') img_byte_arr.seek(0) image = thermal_feet_dataset.load_instance_image(img_byte_arr, merge_image=True) else: image = thermal_feet_dataset.load_instance_image(os.path.join(os.path.dirname(__file__), "assets", "t0.jpg"), merge_image=True) pred_image = predict(model, image) if om: result_image = pred_image.squeeze(2).numpy() else: result_image = show_result(image.numpy(), pred_image, show=False) return result_image def main(): dino_model, sam_model = load_model_def() model = GSamnet(dino_args={"model": dino_model}, sam_args={"model": sam_model, "points": True, "torch": True}).to(sam_model.device) model.eval() model.dummy_input() model_name = "sam2_l" with gr.Blocks() as app: gr.Markdown("### GSAMnet") with gr.Row(): with gr.Column(): uploaded_image = gr.Image(label="Choose a file",value=examples[0],type='pil',image_mode="L") # Cambiado a gr.Image() with gr.Column(): pred_output = gr.Image(label="Prediction",interactive=False,image_mode='L') with gr.Accordion("Advanced options", open=False): om_checkbox = gr.Checkbox(label="Only mask") model_dropdown = gr.Dropdown( label="**Model**", choices=list(SAM2_MODELS.keys()) + list(SAM1_MODELS.keys()), value="sam2_l" ) segment_button = gr.Button("Segment!",variant="primary") # Define the callback for the button def segment_image(uploaded_image, model_dropdown, om): nonlocal model, sam_model, model_name if model_dropdown != model_name: model = update_model(model_dropdown, dino_model) model_name = model_dropdown result = process_image(uploaded_image, model, om=om) return result segment_button.click(segment_image, inputs=[uploaded_image,model_dropdown, om_checkbox], outputs=[pred_output]) gr.Examples(examples=examples, inputs=uploaded_image, label="Ejemplos de imágenes") app.launch() if __name__ == "__main__": main()