CoolFace
Apppublic

DannyLeo/rmbg_api

sourceHugging Facemitupdated 2y agoView on Hugging Face
0likes
app.py133 linesDownload Raw Back to root
1from os import getenv2from flask import Flask, render_template_string, jsonify, request, Response3from pathlib import Path4from werkzeug.utils import secure_filename5from re import findall6from rembg import remove, new_session7from requests import get8from urllib.parse import urlparse9 10app = Flask(__name__)11 12models_path = Path(getenv('U2NET_HOME', './models '))13models_path.mkdir(parents=True, exist_ok=True)14models = ["birefnet-cod", "birefnet-dis", "birefnet-general", "birefnet-general-lite", "birefnet-hrsod", "birefnet-massive", "birefnet-portrait", "isnet-anime", "isnet-general-use", "silueta", "u2net", "u2net_cloth_seg", "u2net_human_seg", "u2netp"]15# models = [file.name for file in models_path.iterdir() if file.is_file() and file.suffix == '.onnx']16 17session = new_session(getenv('DEFAULT_MODEL', 'u2netp'))18 19def hex_to_rgba(hex_code):20    hex_code = hex_code.lstrip('#')21    22    if len(hex_code) not in (6, 8):23        return24        # raise ValueError("Invalid hex code, hex code must be 6 to 9 characters long.")25    26    # Extract the red, green, and blue components27    r = int(hex_code[0:2], 16)  # Convert the first two characters to decimal28    g = int(hex_code[2:4], 16)  # Convert the next two characters to decimal29    b = int(hex_code[4:6], 16)  # Convert the last two characters to decimal30    31    if len(hex_code) == 8:32        a = int(hex_code[6:8], 16)33        return (r, g, b, a)34    35    return (r, g, b, 255)36 37 38 39@app.route('/')40def home():41    html_content = '''42    <!DOCTYPE html>43    <html lang="en">44    <head>45        <meta charset="UTF-8">46        <meta name="viewport" content="width=device-width, initial-scale=1.0">47        <title>Welcome to rmbg_api</title>48        <style>49            body {50                font-family: Arial, sans-serif;51                text-align: center;52                padding: 50px;53            }54            h1 {55                color: #333;56            }57            p {58                font-size: 18px;59            }60        </style>61    </head>62    <body>63        <h1>Welcome to my rmbg_api Application!</h1>64        <p>This api is based on the rembg python package</p>65        <p>Please head to our project GitHub for documentation:</p>66        <a href="https://github.com/DannyAkintunde/rmbg_web" target="_blank">Project GitHub</a><br/>67        <small>clone this space to use the api</small>68    </body>69    </html>70    '''71    return render_template_string(html_content)72 73@app.route('/api/rmbg', methods=['GET', 'POST'])74def remove_bg():75      global session76      try:77          apikey = getenv('APIKEY')78          if apikey and apikey != request.headers.get('X-API-KEY'):79              return jsonify({'error': 'Invalid apikey'}), 40180          if request.method == 'POST':81              file = request.files.get('file')82              model = request.form.get('model')83              bg_color = request.form.get('bg_color')84              85              if not file:86                  return jsonify({'error': 'no file in request'}), 40087              if model:88                  if model not in models:89                      return jsonify({'error': f'invalid model name avaliable models are {models}'}), 40090                  session = new_session(model)91              92              filename = secure_filename(file.filename)93              mime_type = file.content_type94              95              data = file.stream.read()96          elif request.method == 'GET':97              url = request.args.get('url')98              model = request.args.get('model')99              bg_color = request.args.get('bg_color')100              101              if not url:102                  return jsonify({'error': 'Url is required'}), 400103              104              parsed_url = urlparse(url)105              if not parsed_url.scheme or not parsed_url.netloc:106                  return jsonify({'error': 'Invalid url passed'}), 400107              108              if model:109                  if model not in models:110                      return jsonify({'error': f'invalid model name avaliable models are {models}'}), 400111                  session = new_session(model)112              113              input_image_response = get(url)114              input_image_response.raise_for_status()115              if input_image_response.status_code == 200:116                  data = input_image_response.content117                  content_disposition = input_image_response.headers.get('Content-Disposition', '')118                  filename_search = findall("filename=(.+)", content_disposition)119                  filename = filename_search[0].strip('"').strip("'") if filename_search else url.split("/")[-1]120                  mime_type = input_image_response.headers.get('Content-Type', 'application/octet-stream')121                  122          bg_color = hex_to_rgba(bg_color) if bg_color else None123          result = remove(data, bgcolor=bg_color, session=session, alpha_matting=True, alpha_matting_foreground_threshold=270, alpha_matting_background_threshold=20, alpha_matting_erode_size=11, post_process_mask=True)124          return Response(response=result, mimetype=mime_type, headers={125                'Content-Disposition': f'inline; filename=removedbg_{filename}'126            })127      except Exception as e:128          print(e)129          return jsonify({'error': f'Server error {str(e)}'}), 500130  131if __name__ == '__main__':132    app.run(host=getenv('HOST', '127.0.0.1'),133            port=getenv('PORT', 8000), debug=(getenv('DEBUG', '').lower() == 'true'))