返回提交历史
Modified
g4f/Provider/nexra/NexraSDTurbo.py
+41
-40
Modified
g4f/models.py
+10
-1
XFEstudio/gpt4free
Restored provider (g4f/Provider/nexra/NexraSDTurbo.py)
b08249ec
代码差异
2 个文件
+51
-41
@@ -1,28 +1,26 @@
1
1
from __future__ import annotations
2
2
3
from aiohttp import ClientSession
4
3
import json
5
6
from ...typing import AsyncResult, Messages
7
from ..base_provider import AsyncGeneratorProvider, ProviderModelMixin
4
import requests
5
from ...typing import CreateResult, Messages
6
from ..base_provider import ProviderModelMixin, AbstractProvider
8
7
from ...image import ImageResponse
9
8
10
11
class NexraSDTurbo(AsyncGeneratorProvider, ProviderModelMixin):
9
class NexraSDTurbo(AbstractProvider, ProviderModelMixin):
12
10
label = "Nexra Stable Diffusion Turbo"
13
11
url = "https://nexra.aryahcr.cc/documentation/stable-diffusion/en"
14
12
api_endpoint = "https://nexra.aryahcr.cc/api/image/complements"
15
working = False
13
working = True
16
14
17
default_model = 'sdxl-turbo'
15
default_model = "sdxl-turbo"
18
16
models = [default_model]
19
17
20
18
@classmethod
21
19
def get_model(cls, model: str) -> str:
22
20
return cls.default_model
23
21
24
22
@classmethod
25
async def create_async_generator(
23
def create_completion(
26
24
cls,
27
25
model: str,
28
26
messages: Messages,
@@ -31,38 +29,41 @@ class NexraSDTurbo(AsyncGeneratorProvider, ProviderModelMixin):
31
29
strength: str = 0.7, # Min: 0, Max: 1
32
30
steps: str = 2, # Min: 1, Max: 10
33
31
**kwargs
34
) -> AsyncResult:
32
) -> CreateResult:
35
33
model = cls.get_model(model)
36
34
37
35
headers = {
38
"Content-Type": "application/json"
36
'Content-Type': 'application/json'
39
37
}
40
async with ClientSession(headers=headers) as session:
41
prompt = messages[0]['content']
42
data = {
43
"prompt": prompt,
44
"model": model,
45
"response": response,
46
"data": {
47
"strength": strength,
48
"steps": steps
49
}
38
39
data = {
40
"prompt": messages[-1]["content"],
41
"model": model,
42
"response": response,
43
"data": {
44
"strength": strength,
45
"steps": steps
50
46
}
51
async with session.post(cls.api_endpoint, json=data, proxy=proxy) as response:
52
text_data = await response.text()
53
54
if response.status == 200:
55
try:
56
json_start = text_data.find('{')
57
json_data = text_data[json_start:]
58
59
data = json.loads(json_data)
60
if 'images' in data and len(data['images']) > 0:
61
image_url = data['images'][-1]
62
yield ImageResponse(image_url, prompt)
63
else:
64
yield ImageResponse("No images found in the response.", prompt)
65
except json.JSONDecodeError:
66
yield ImageResponse("Failed to parse JSON. Response might not be in JSON format.", prompt)
47
}
48
49
response = requests.post(cls.api_endpoint, headers=headers, json=data)
50
51
result = cls.process_response(response)
52
yield result
53
54
@classmethod
55
def process_response(cls, response):
56
if response.status_code == 200:
57
try:
58
content = response.text.strip()
59
content = content.lstrip('_') # Remove the leading underscore
60
data = json.loads(content)
61
if data.get('status') and data.get('images'):
62
image_url = data['images'][0]
63
return ImageResponse(images=[image_url], alt="Generated Image")
67
64
else:
68
yield ImageResponse(f"Request failed with status: {response.status}", prompt)
65
return "Error: No image URL found in the response"
66
except json.JSONDecodeError as e:
67
return f"Error: Unable to decode JSON response. Details: {str(e)}"
68
else:
69
return f"Error: {response.status_code}, Response: {response.text}"
@@ -53,6 +53,7 @@ from .Provider import (
53
53
NexraMidjourney,
54
54
NexraQwen,
55
55
NexraSD15,
56
NexraSDTurbo,
56
57
OpenaiChat,
57
58
PerplexityLabs,
58
59
Pi,
@@ -734,10 +735,17 @@ nemotron_70b = Model(
734
735
#############
735
736
736
737
### Stability AI ###
738
sdxl_turbo = Model(
739
name = 'sdxl-turbo',
740
base_provider = 'Stability AI',
741
best_provider = NexraSDTurbo
742
743
)
744
737
745
sdxl = Model(
738
746
name = 'sdxl',
739
747
base_provider = 'Stability AI',
740
best_provider = IterListProvider([ReplicateHome, DeepInfraImage])
748
best_provider = IterListProvider([ReplicateHome, DeepInfraImage, sdxl_turbo.best_provider])
741
749
742
750
)
743
751
@@ -1103,6 +1111,7 @@ class ModelUtils:
1103
1111
1104
1112
### Stability AI ###
1105
1113
'sdxl': sdxl,
1114
'sdxl-turbo': sdxl_turbo,
1106
1115
'sd-1.5': sd_1_5,
1107
1116
'sd-3': sd_3,
1108
1117