File size: 1,544 Bytes
645fa0f
 
67902ab
645fa0f
67902ab
645fa0f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
f35cb88
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
import os
import io
# import IPython.display
from PIL import Image
# import base64 
import re

from dotenv import load_dotenv, find_dotenv
_ = load_dotenv(find_dotenv()) # read local .env file
# hf_api_key = os.environ['HF_API_KEY']
hf_api_key = os.getenv('HF_API_KEY')
HF_API_TTI_BASE='https://api-inference.huggingface.co/models/runwayml/stable-diffusion-v1-5'

# Helper function
import requests, json

#Text-to-image endpoint
def get_completion(inputs, parameters=None, ENDPOINT_URL=HF_API_TTI_BASE):
    headers = {
      "Authorization": f"Bearer {hf_api_key}",
      "Content-Type": "application/json"
    }   
    data = { "inputs": inputs }
    if parameters is not None:
        data.update({"parameters": parameters})
    response = requests.request("POST",
                                ENDPOINT_URL,
                                headers=headers,
                                data=json.dumps(data))
    image_bytes = response.content
    return image_bytes

def base64_to_pil(image_bytes):
    # base64_decoded = base64.b64decode(img_base64)
    byte_stream = io.BytesIO(image_bytes)
    byte_stream.seek(0)
    try:
        pil_image = Image.open(byte_stream)
        pil_image.verify()
    except (IOError, SyntaxError) as e:
        print("Error: ", e)
    pil_image = Image.open(byte_stream)
    return pil_image

def generate(prompt):
    prompt = re.sub(' +', ' ', prompt)
    prompt = re.sub('[^A-Za-z0-9., ]+', '', prompt)
    output = get_completion(prompt)
    pil_image = base64_to_pil(output)
    return pil_image