horsey-defy commited on
Commit
20e4ac4
·
1 Parent(s): 9a46212

Added list_models, list_samplers and wait, fixed Org strings

Browse files
Files changed (1) hide show
  1. app.py +25 -9
app.py CHANGED
@@ -16,18 +16,34 @@ class RenderNet:
16
  def __init__(self, api_key, base=None):
17
  self.base = base or "https://api.prodia.com/v1"
18
  self.headers = {
19
- "X-Prodia-Key": api_key
20
  }
21
 
22
  def generate(self, params):
23
  response = self._post(f"{self.base}/sd/generate", params)
24
  return response.json()
25
 
 
 
 
 
 
 
 
 
 
 
 
 
26
 
27
  def list_models(self):
28
  response = self._get(f"{self.base}/sd/models")
29
  return response.json()
30
 
 
 
 
 
31
  def _post(self, url, params):
32
  headers = {
33
  **self.headers,
@@ -36,7 +52,7 @@ class RenderNet:
36
  response = requests.post(url, headers=headers, data=json.dumps(params))
37
 
38
  if response.status_code != 200:
39
- raise Exception(f"Bad Prodia Response: {response.status_code}")
40
 
41
  return response
42
 
@@ -44,7 +60,7 @@ class RenderNet:
44
  response = requests.get(url, headers=self.headers)
45
 
46
  if response.status_code != 200:
47
- raise Exception(f"Bad Prodia Response: {response.status_code}")
48
 
49
  return response
50
 
@@ -131,8 +147,8 @@ def send_to_txt2img(image):
131
  return result
132
 
133
 
134
- prodia_client = Prodia(api_key=os.getenv("PRODIA_API_KEY"))
135
- model_list = prodia_client.list_models()
136
  model_names = {}
137
 
138
  for model_name in model_list:
@@ -140,7 +156,7 @@ for model_name in model_list:
140
  model_names[name_without_ext] = model_name
141
 
142
  def txt2img(prompt, negative_prompt, model, steps, sampler, cfg_scale, width, height, seed):
143
- result = prodia_client.generate({
144
  "prompt": prompt,
145
  "negative_prompt": negative_prompt,
146
  "model": model,
@@ -152,7 +168,7 @@ def txt2img(prompt, negative_prompt, model, steps, sampler, cfg_scale, width, he
152
  "seed": seed
153
  })
154
 
155
- job = prodia_client.wait(result)
156
 
157
  return job["imageUrl"]
158
 
@@ -166,10 +182,10 @@ css = """
166
  with gr.Blocks(css=css) as demo:
167
  with gr.Row():
168
  with gr.Column(scale=6):
169
- model = gr.Dropdown(interactive=True,value="absolutereality_v181.safetensors [3d9d4d2b]", show_label=True, label="Stable Diffusion Checkpoint", choices=prodia_client.list_models())
170
 
171
  with gr.Column(scale=1):
172
- gr.Markdown(elem_id="powered-by-prodia", value="AUTOMATIC1111 Stable Diffusion Web UI.<br>Powered by [Prodia](https://prodia.com).<br>For more features and faster generation times check out our [API Docs](https://docs.prodia.com/reference/getting-started-guide).")
173
 
174
 
175
  with gr.Tabs() as tabs:
 
16
  def __init__(self, api_key, base=None):
17
  self.base = base or "https://api.prodia.com/v1"
18
  self.headers = {
19
+ "X-RenderNet-Key": api_key
20
  }
21
 
22
  def generate(self, params):
23
  response = self._post(f"{self.base}/sd/generate", params)
24
  return response.json()
25
 
26
+ def get_job(self, job_id):
27
+ response = self._get(f"{self.base}/job/{job_id}")
28
+ return response.json()
29
+
30
+ def wait(self, job):
31
+ job_result = job
32
+
33
+ while job_result['status'] not in ['succeeded', 'failed']:
34
+ time.sleep(0.25)
35
+ job_result = self.get_job(job['job'])
36
+
37
+ return job_result
38
 
39
  def list_models(self):
40
  response = self._get(f"{self.base}/sd/models")
41
  return response.json()
42
 
43
+ def list_samplers(self):
44
+ response = self._get(f"{self.base}/sd/samplers")
45
+ return response.json()
46
+
47
  def _post(self, url, params):
48
  headers = {
49
  **self.headers,
 
52
  response = requests.post(url, headers=headers, data=json.dumps(params))
53
 
54
  if response.status_code != 200:
55
+ raise Exception(f"Bad RenderNet Response: {response.status_code}")
56
 
57
  return response
58
 
 
60
  response = requests.get(url, headers=self.headers)
61
 
62
  if response.status_code != 200:
63
+ raise Exception(f"Bad RenderNet Response: {response.status_code}")
64
 
65
  return response
66
 
 
147
  return result
148
 
149
 
150
+ rendernet_client = RenderNet(api_key=os.getenv("RENDERNET_API_KEY"))
151
+ model_list = rendernet_client.list_models()
152
  model_names = {}
153
 
154
  for model_name in model_list:
 
156
  model_names[name_without_ext] = model_name
157
 
158
  def txt2img(prompt, negative_prompt, model, steps, sampler, cfg_scale, width, height, seed):
159
+ result = rendernet_client.generate({
160
  "prompt": prompt,
161
  "negative_prompt": negative_prompt,
162
  "model": model,
 
168
  "seed": seed
169
  })
170
 
171
+ job = rendernet_client.wait(result)
172
 
173
  return job["imageUrl"]
174
 
 
182
  with gr.Blocks(css=css) as demo:
183
  with gr.Row():
184
  with gr.Column(scale=6):
185
+ model = gr.Dropdown(interactive=True,value="absolutereality_v181.safetensors [3d9d4d2b]", show_label=True, label="Stable Diffusion Checkpoint", choices=rendernet_client.list_models())
186
 
187
  with gr.Column(scale=1):
188
+ gr.Markdown(elem_id="powered-by-rendernet", value="AUTOMATIC1111 Stable Diffusion Web UI.<br>Powered by [RenderNet](https://rendernet.ai).<br>For more features and faster generation times check out our [API Docs](https://rendernet.ai/).")
189
 
190
 
191
  with gr.Tabs() as tabs: