XFE Git
XFE Studio Git
Git 首页 全局搜索
XFE 主站 文档 NuGet
公开
关注 0 Fork 0 Star 1
返回提交历史

XFEstudio/gpt4free

Restored provider (g4f/Provider/nexra/NexraSDLora.py)

5a79d8cb
kqlio67 <kqlio67@users.noreply.github.com>
提交于

代码差异

2 个文件 +51 -41
Modified g4f/Provider/nexra/NexraSDLora.py +41 -40
@@ -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 NexraSDLora(AsyncGeneratorProvider, ProviderModelMixin):
9 class NexraSDLora(AbstractProvider, ProviderModelMixin):
12 10 label = "Nexra Stable Diffusion Lora"
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-lora'
15 default_model = "sdxl-lora"
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 NexraSDLora(AsyncGeneratorProvider, ProviderModelMixin):
31 29 guidance: str = 0.3, # Min: 0, Max: 5
32 30 steps: str = 2, # Min: 2, 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 "guidance": guidance,
48 "steps": steps
49 }
38
39 data = {
40 "prompt": messages[-1]["content"],
41 "model": model,
42 "response": response,
43 "data": {
44 "guidance": guidance,
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('_')
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}"
Modified g4f/models.py +10 -1
@@ -53,6 +53,7 @@ from .Provider import (
53 53 NexraMidjourney,
54 54 NexraQwen,
55 55 NexraSD15,
56 NexraSDLora,
56 57 NexraSDTurbo,
57 58 OpenaiChat,
58 59 PerplexityLabs,
@@ -742,10 +743,17 @@ sdxl_turbo = Model(
742 743
743 744 )
744 745
746 sdxl_lora = Model(
747 name = 'sdxl-lora',
748 base_provider = 'Stability AI',
749 best_provider = NexraSDLora
750
751 )
752
745 753 sdxl = Model(
746 754 name = 'sdxl',
747 755 base_provider = 'Stability AI',
748 best_provider = IterListProvider([ReplicateHome, DeepInfraImage, sdxl_turbo.best_provider])
756 best_provider = IterListProvider([ReplicateHome, DeepInfraImage, sdxl_turbo.best_provider, sdxl_lora.best_provider])
749 757
750 758 )
751 759
@@ -1111,6 +1119,7 @@ class ModelUtils:
1111 1119
1112 1120 ### Stability AI ###
1113 1121 'sdxl': sdxl,
1122 'sdxl-lora': sdxl_lora,
1114 1123 'sdxl-turbo': sdxl_turbo,
1115 1124 'sd-1.5': sd_1_5,
1116 1125 'sd-3': sd_3,