| from flask import Flask, render_template, request, jsonify, send_file |
| import altair as alt |
| import numpy as np |
| import pandas as pd |
| import json |
| import os |
| import tempfile |
| from werkzeug.utils import secure_filename |
| import torch |
| from PIL import Image |
| import cv2 |
| import io |
| import uuid |
|
|
| |
| from groundingdino.models import build_model |
| from groundingdino.util.slconfig import SLConfig |
| from groundingdino.util.utils import clean_state_dict |
| from groundingdino.util.inference import load_image, predict, annotate |
|
|
| |
| CONFIG_PATH = "/app/GroundingDINO/groundingdino/config/GroundingDINO_SwinT_OGC.py" |
| CHECKPOINT_PATH = "/app/weights/groundingdino_swint_ogc.pth" |
|
|
| app = Flask(__name__) |
| app.config['MAX_CONTENT_LENGTH'] = 16 * 1024 * 1024 |
|
|
| |
| UPLOAD_FOLDER = '/tmp/uploads' |
| if not os.path.exists(UPLOAD_FOLDER): |
| os.makedirs(UPLOAD_FOLDER) |
| |
| |
| label_data_store = {} |
|
|
| |
| ALLOWED_EXTENSIONS = {'png', 'jpg', 'jpeg'} |
|
|
| def allowed_file(filename): |
| return '.' in filename and filename.rsplit('.', 1)[1].lower() in ALLOWED_EXTENSIONS |
|
|
| def load_grounding_dino_model(): |
| """Load the GroundingDINO model""" |
| try: |
| args = SLConfig.fromfile(CONFIG_PATH) |
| model = build_model(args) |
| checkpoint = torch.load(CHECKPOINT_PATH, map_location="cpu") |
| model.load_state_dict(clean_state_dict(checkpoint["model"]), strict=False) |
| model.eval() |
| return model |
| except Exception as e: |
| print(f"Error loading model: {e}") |
| return None |
|
|
| |
| model = load_grounding_dino_model() |
|
|
| @app.route('/') |
| def index(): |
| return render_template('index.html') |
|
|
| @app.route('/health') |
| def health(): |
| return {'status': 'healthy'} |
|
|
| @app.route('/label_image', methods=['POST']) |
| def label_image(): |
| """Endpoint to label images with GroundingDINO using text prompts""" |
| if model is None: |
| return jsonify({'error': 'Model not loaded'}), 500 |
| |
| try: |
| |
| if 'image' not in request.files: |
| return jsonify({'error': 'No image file provided'}), 400 |
| |
| file = request.files['image'] |
| if file.filename == '': |
| return jsonify({'error': 'No image selected'}), 400 |
| |
| if not allowed_file(file.filename): |
| return jsonify({'error': 'Invalid file type'}), 400 |
| |
| |
| text_prompt = request.form.get('text_prompt', '') |
| if not text_prompt: |
| return jsonify({'error': 'No text prompt provided'}), 400 |
| |
| |
| box_threshold = float(request.form.get('box_threshold', 0.35)) |
| text_threshold = float(request.form.get('text_threshold', 0.25)) |
| |
| |
| filename = secure_filename(file.filename) |
| filepath = os.path.join(UPLOAD_FOLDER, filename) |
| file.save(filepath) |
| |
| |
| image_source, image = load_image(filepath) |
| |
| |
| image = image.cpu() if hasattr(image, 'is_cuda') and image.is_cuda else image |
| |
| boxes, logits, phrases = predict( |
| model=model, |
| image=image, |
| caption=text_prompt, |
| box_threshold=box_threshold, |
| text_threshold=text_threshold, |
| device="cpu" |
| ) |
| |
| |
| annotated_frame = annotate( |
| image_source=image_source, |
| boxes=boxes, |
| logits=logits, |
| phrases=phrases |
| ) |
| |
| |
| output_path = os.path.join(UPLOAD_FOLDER, f"labeled_{filename}") |
| cv2.imwrite(output_path, annotated_frame) |
| |
| |
| image_id = str(uuid.uuid4()) |
| label_data_store[image_id] = { |
| 'filename': filename, |
| 'original_path': filepath, |
| 'labeled_path': output_path, |
| 'boxes': boxes.tolist() if hasattr(boxes, 'tolist') else [], |
| 'logits': logits.tolist() if hasattr(logits, 'tolist') else [], |
| 'phrases': phrases, |
| 'text_prompt': text_prompt, |
| 'box_threshold': box_threshold, |
| 'text_threshold': text_threshold |
| } |
| |
| |
| label_data_store[filename] = label_data_store[image_id] |
| |
| |
| return jsonify({ |
| 'message': 'Image processed successfully', |
| 'image_id': image_id, |
| 'boxes': boxes.tolist() if hasattr(boxes, 'tolist') else [], |
| 'logits': logits.tolist() if hasattr(logits, 'tolist') else [], |
| 'phrases': phrases, |
| 'original_image_url': f'/original_image/{filename}', |
| 'labeled_image_url': f'/labeled_image/{filename}' |
| }) |
| |
| except Exception as e: |
| return jsonify({'error': str(e)}), 500 |
|
|
| @app.route('/labeled_image/<filename>') |
| def get_labeled_image(filename): |
| """Serve labeled images""" |
| try: |
| labeled_filepath = os.path.join(UPLOAD_FOLDER, f"labeled_{filename}") |
| if os.path.exists(labeled_filepath): |
| return send_file(labeled_filepath, mimetype='image/jpeg') |
| else: |
| return jsonify({'error': 'Labeled image not found'}), 404 |
| except Exception as e: |
| return jsonify({'error': str(e)}), 500 |
|
|
| @app.route('/original_image/<filename>') |
| def get_original_image(filename): |
| """Serve original images""" |
| try: |
| filepath = os.path.join(UPLOAD_FOLDER, filename) |
| if os.path.exists(filepath): |
| return send_file(filepath, mimetype='image/jpeg') |
| else: |
| return jsonify({'error': 'Original image not found'}), 404 |
| except Exception as e: |
| return jsonify({'error': str(e)}), 500 |
|
|
| @app.route('/image_blob/<identifier>') |
| def get_image_blob(identifier): |
| """Endpoint to get image blob directly by filename or image ID""" |
| try: |
| |
| if os.path.exists(os.path.join(UPLOAD_FOLDER, identifier)): |
| filepath = os.path.join(UPLOAD_FOLDER, identifier) |
| return send_file(filepath, mimetype='image/jpeg') |
| |
| |
| labeled_filepath = os.path.join(UPLOAD_FOLDER, f"labeled_{identifier}") |
| if os.path.exists(labeled_filepath): |
| return send_file(labeled_filepath, mimetype='image/jpeg') |
| |
| |
| if identifier in label_data_store: |
| image_data = label_data_store[identifier] |
| labeled_path = image_data.get('labeled_path') |
| if labeled_path and os.path.exists(labeled_path): |
| return send_file(labeled_path, mimetype='image/jpeg') |
| elif os.path.exists(image_data.get('original_path', '')): |
| return send_file(image_data['original_path'], mimetype='image/jpeg') |
| |
| return jsonify({'error': 'Image not found'}), 404 |
| except Exception as e: |
| return jsonify({'error': str(e)}), 500 |
|
|
| @app.route('/labels_json/<identifier>') |
| def get_labels_json(identifier): |
| """Endpoint to get labels as JSON by filename or image ID""" |
| try: |
| |
| if identifier in label_data_store: |
| image_data = label_data_store[identifier] |
| return jsonify({ |
| 'filename': image_data['filename'], |
| 'boxes': image_data['boxes'], |
| 'logits': image_data['logits'], |
| 'phrases': image_data['phrases'], |
| 'text_prompt': image_data['text_prompt'], |
| 'box_threshold': image_data['box_threshold'], |
| 'text_threshold': image_data['text_threshold'] |
| }) |
| |
| |
| for image_id, image_data in label_data_store.items(): |
| if image_data['filename'] == identifier: |
| return jsonify({ |
| 'filename': image_data['filename'], |
| 'boxes': image_data['boxes'], |
| 'logits': image_data['logits'], |
| 'phrases': image_data['phrases'], |
| 'text_prompt': image_data['text_prompt'], |
| 'box_threshold': image_data['box_threshold'], |
| 'text_threshold': image_data['text_threshold'] |
| }) |
| |
| return jsonify({'error': 'No labels found for this image'}), 404 |
| except Exception as e: |
| return jsonify({'error': str(e)}), 500 |
|
|
| if __name__ == '__main__': |
| app.run(host='0.0.0.0', port=8501, debug=True) |