-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathhuggingface_api.py
More file actions
153 lines (127 loc) · 5.19 KB
/
Copy pathhuggingface_api.py
File metadata and controls
153 lines (127 loc) · 5.19 KB
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
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
import os
import requests
import json
from PIL import Image
import io
import base64
from utils import image_to_base64, base64_to_image
class HuggingFaceAPI:
"""Service to interact with Hugging Face Inference API"""
def __init__(self, model_id="CompVis/stable-diffusion-v1-4", api_token=None):
# Try getting token from parameter first
self.api_token = api_token
# If not provided, try environment variable
if not self.api_token:
self.api_token = os.environ.get("HF_API_TOKEN")
# If still not found, try reading from config file
if not self.api_token:
try:
if os.path.exists('config.json'):
with open('config.json', 'r') as f:
config = json.load(f)
self.api_token = config.get('hf_api_token')
except:
pass
# If still not found, raise error
if not self.api_token:
raise ValueError(
"Hugging Face API token not found. Please provide it through one of these methods:\n"
"1. Pass directly to the HuggingFaceAPI constructor\n"
"2. Set the HF_API_TOKEN environment variable\n"
"3. Create a config.json file with an 'hf_api_token' field"
)
self.model_id = model_id
self.api_url = f"https://api-inference.huggingface.co/models/{model_id}"
self.headers = {"Authorization": f"Bearer {self.api_token}"}
def generate_image(
self,
prompt,
negative_prompt="",
height=512,
width=512,
num_inference_steps=50,
guidance_scale=7.5,
seed=None,
):
"""Generate an image from a text prompt using HF API"""
payload = {
"inputs": prompt,
"parameters": {
"negative_prompt": negative_prompt,
"height": height,
"width": width,
"num_inference_steps": num_inference_steps,
"guidance_scale": guidance_scale,
}
}
# Add seed if provided
if seed is not None:
payload["parameters"]["seed"] = seed
# Make API request
response = requests.post(self.api_url, headers=self.headers, json=payload)
if response.status_code != 200:
raise Exception(f"API request failed with status code {response.status_code}: {response.text}")
# The response is the binary image data
image = Image.open(io.BytesIO(response.content))
# Convert to base64 for API response
base64_image = image_to_base64(image)
# Get the seed used (if provided in response, otherwise use input seed or None)
result_seed = seed
if "seed" in response.headers:
result_seed = int(response.headers.get("seed"))
return {
"image": base64_image,
"prompt": prompt,
"seed": result_seed,
}
def generate_variations(
self,
image,
prompt="",
negative_prompt="",
strength=0.75,
num_inference_steps=50,
guidance_scale=7.5,
num_variations=4,
):
"""Generate variations of an input image using img2img"""
# Convert PIL Image to bytes
img_byte_arr = io.BytesIO()
image.save(img_byte_arr, format='PNG')
img_byte_arr = img_byte_arr.getvalue()
# For img2img we need to use a different endpoint
img2img_url = f"https://api-inference.huggingface.co/models/{self.model_id}/img2img"
variations = []
# Generate multiple variations
for _ in range(num_variations):
files = {
'image': img_byte_arr,
}
data = {
'prompt': prompt,
'negative_prompt': negative_prompt,
'strength': strength,
'guidance_scale': guidance_scale,
'num_inference_steps': num_inference_steps,
}
# Make API request
response = requests.post(img2img_url, headers=self.headers, files=files, data=data)
if response.status_code != 200:
raise Exception(f"API request failed with status code {response.status_code}: {response.text}")
# The response is the binary image data
variation_image = Image.open(io.BytesIO(response.content))
# Convert to base64 for API response
base64_image = image_to_base64(variation_image)
# Get seed if available in response headers, otherwise generate a random one
seed = None
if "seed" in response.headers:
seed = int(response.headers.get("seed"))
else:
import random
seed = random.randint(0, 2**32 - 1)
variations.append({
"image": base64_image,
"prompt": prompt,
"seed": seed,
})
return variations