DeepShape_Space / TextToImage_StableDiffusion.py
arvindershinh's picture
commented import packages
67902ab verified
Raw
History Blame Contribute Delete
1.54 kB
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