Arulkumar03/Wheat_HEAD_Detection_Counting_ComputerVision_Model
0
1import cv22import numpy as np3import gradio as gr4from detectron2 import model_zoo5from detectron2.config import get_cfg6from detectron2.engine import DefaultPredictor7from detectron2.utils.visualizer import Visualizer8from detectron2.data import MetadataCatalog9 10def initialize_model():11 for d in ["train", "test"]:12 #DatasetCatalog.register("wheat_" + d, lambda d=d: get_wheat_dicts("wheat_Detection/" + d))13 MetadataCatalog.get("wheat_" + d).set(thing_classes=["wheat"])14 15 wheat_metadata = MetadataCatalog.get("wheat_train") 16 cfg = get_cfg()17 cfg.MODEL.DEVICE = "cpu"18 cfg.DATALOADER.NUM_WORKERS = 019 cfg.MODEL.WEIGHTS = model_zoo.get_checkpoint_url("COCO-InstanceSegmentation/mask_rcnn_R_101_C4_3x.yaml")20 cfg.SOLVER.IMS_PER_BATCH = 221 cfg.SOLVER.BASE_LR = 0.0002522 cfg.SOLVER.STEPS = []23 cfg.MODEL.ROI_HEADS.BATCH_SIZE_PER_IMAGE = 12824 cfg.MODEL.ROI_HEADS.NUM_CLASSES = 225 cfg.MODEL.WEIGHTS = "output/model_final.pth"26 cfg.MODEL.ROI_HEADS.SCORE_THRESH_TEST = 0.9527 predictor = DefaultPredictor(cfg)28 return predictor29 30def process_image(predictor, img):31 outputs = predictor(img)32 wheat_metadata = MetadataCatalog.get("wheat_train")33 v = Visualizer(img[:, :, ::-1],34 metadata=wheat_metadata, 35 scale=1.5, 36 instance_mode="segmentation")37 out = v.draw_instance_predictions(outputs["instances"].to("cpu"))38 processed_img = cv2.cvtColor(out.get_image()[:, :, ::-1], cv2.COLOR_BGR2RGB)39 return processed_img40 41def main(img):42 predictor = initialize_model()43 processed_img = process_image(predictor, img)44 return processed_img45 46 47iface = gr.Interface(48 fn=main,49 inputs="image",50 outputs="image",51 title="Wheat head Detector & Counting Wheat heads",52 cache_examples=False, port=7861).launch(share=True)