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

XFEstudio/gpt4free

Add more flux dev image providers

1bfb3617
Heiner Lohaus <hlohaus@users.noreply.github.com>
提交于

代码差异

9 个文件 +132 -39
Added g4f/Provider/Flux.py +58 -0
@@ -0,0 +1,58 @@
1 from __future__ import annotations
2
3 import json
4 from aiohttp import ClientSession
5
6 from ..typing import AsyncResult, Messages
7 from ..image import ImageResponse, ImagePreview
8 from .base_provider import AsyncGeneratorProvider, ProviderModelMixin
9
10 class Flux(AsyncGeneratorProvider, ProviderModelMixin):
11 label = "Flux Provider"
12 url = "https://black-forest-labs-flux-1-dev.hf.space"
13 api_endpoint = "/gradio_api/call/infer"
14 working = True
15 default_model = 'flux-1-dev'
16 models = [default_model]
17 image_models = [default_model]
18
19 @classmethod
20 async def create_async_generator(
21 cls, model: str, messages: Messages, prompt: str = None, api_key: str = None, proxy: str = None, **kwargs
22 ) -> AsyncResult:
23 headers = {
24 "Content-Type": "application/json",
25 "Accept": "application/json",
26 }
27 if api_key is not None:
28 headers["Authorization"] = f"Bearer {api_key}"
29 async with ClientSession(headers=headers) as session:
30 prompt = messages[-1]["content"] if prompt is None else prompt
31 data = {
32 "data": [prompt, 0, True, 1024, 1024, 3.5, 28]
33 }
34 async with session.post(f"{cls.url}{cls.api_endpoint}", json=data, proxy=proxy) as response:
35 response.raise_for_status()
36 event_id = (await response.json()).get("event_id")
37 async with session.get(f"{cls.url}{cls.api_endpoint}/{event_id}") as event_response:
38 event_response.raise_for_status()
39 event = None
40 async for chunk in event_response.content:
41 if chunk.startswith(b"event: "):
42 event = chunk[7:].decode(errors="replace").strip()
43 if chunk.startswith(b"data: "):
44 if event == "error":
45 raise RuntimeError(f"GPU token limit exceeded: {chunk.decode(errors='replace')}")
46 if event in ("complete", "generating"):
47 try:
48 data = json.loads(chunk[6:])
49 if data is None:
50 continue
51 url = data[0]["url"]
52 except (json.JSONDecodeError, KeyError, TypeError) as e:
53 raise RuntimeError(f"Failed to parse image URL: {chunk.decode(errors='replace')}", e)
54 if event == "generating":
55 yield ImagePreview(url, prompt)
56 else:
57 yield ImageResponse(url, prompt)
58 break
Modified g4f/Provider/__init__.py +2 -1
@@ -39,6 +39,7 @@ from .TeachAnything import TeachAnything
39 39 from .Upstage import Upstage
40 40 from .You import You
41 41 from .Mhystical import Mhystical
42 from .Flux import Flux
42 43
43 44 import sys
44 45
@@ -59,4 +60,4 @@ __map__: dict[str, ProviderType] = dict([
59 60 ])
60 61
61 62 class ProviderUtils:
62 convert: dict[str, ProviderType] = __map__
63 convert: dict[str, ProviderType] = __map__
Modified g4f/Provider/needs_auth/HuggingChat.py +15 -7
@@ -12,6 +12,7 @@ from ...typing import CreateResult, Messages, Cookies
12 12 from ...errors import MissingRequirementsError
13 13 from ...requests.raise_for_status import raise_for_status
14 14 from ...cookies import get_cookies
15 from ...image import ImageResponse
15 16 from ..base_provider import ProviderModelMixin, AbstractProvider, BaseConversation
16 17 from ..helper import format_prompt
17 18 from ... import debug
@@ -26,10 +27,12 @@ class HuggingChat(AbstractProvider, ProviderModelMixin):
26 27 working = True
27 28 supports_stream = True
28 29 needs_auth = True
29 default_model = "meta-llama/Meta-Llama-3.1-70B-Instruct"
30
30 default_model = "Qwen/Qwen2.5-72B-Instruct"
31 image_models = [
32 "black-forest-labs/FLUX.1-dev"
33 ]
31 34 models = [
32 'Qwen/Qwen2.5-72B-Instruct',
35 default_model,
33 36 'meta-llama/Meta-Llama-3.1-70B-Instruct',
34 37 'CohereForAI/c4ai-command-r-plus-08-2024',
35 38 'Qwen/QwQ-32B-Preview',
@@ -39,8 +42,8 @@ class HuggingChat(AbstractProvider, ProviderModelMixin):
39 42 'NousResearch/Hermes-3-Llama-3.1-8B',
40 43 'mistralai/Mistral-Nemo-Instruct-2407',
41 44 'microsoft/Phi-3.5-mini-instruct',
45 *image_models
42 46 ]
43
44 47 model_aliases = {
45 48 "qwen-2.5-72b": "Qwen/Qwen2.5-72B-Instruct",
46 49 "llama-3.1-70b": "meta-llama/Meta-Llama-3.1-70B-Instruct",
@@ -52,6 +55,7 @@ class HuggingChat(AbstractProvider, ProviderModelMixin):
52 55 "hermes-3": "NousResearch/Hermes-3-Llama-3.1-8B",
53 56 "mistral-nemo": "mistralai/Mistral-Nemo-Instruct-2407",
54 57 "phi-3.5-mini": "microsoft/Phi-3.5-mini-instruct",
58 "flux-dev": "black-forest-labs/FLUX.1-dev",
55 59 }
56 60
57 61 @classmethod
@@ -109,7 +113,7 @@ class HuggingChat(AbstractProvider, ProviderModelMixin):
109 113 "is_retry": False,
110 114 "is_continue": False,
111 115 "web_search": web_search,
112 "tools": []
116 "tools": ["000000000000000000000001"] if model in cls.image_models else [],
113 117 }
114 118
115 119 headers = {
@@ -162,14 +166,18 @@ class HuggingChat(AbstractProvider, ProviderModelMixin):
162 166
163 167 elif line["type"] == "finalAnswer":
164 168 break
165
166 full_response = full_response.replace('<|im_end|', '').replace('\u0000', '').strip()
169 elif line["type"] == "file":
170 url = f"https://huggingface.co/chat/conversation/{conversation.conversation_id}/output/{line['sha']}"
171 yield ImageResponse(url, alt=messages[-1]["content"], options={"cookies": cookies})
167 172
173 full_response = full_response.replace('<|im_end|', '').replace('\u0000', '').strip()
168 174 if not stream:
169 175 yield full_response
170 176
171 177 @classmethod
172 178 def create_conversation(cls, session: Session, model: str):
179 if model in cls.image_models:
180 model = cls.default_model
173 181 json_data = {
174 182 'model': model,
175 183 }
Modified g4f/Provider/needs_auth/HuggingFace.py +27 -11
@@ -1,21 +1,25 @@
1 1 from __future__ import annotations
2 2
3 3 import json
4 import base64
5 import random
4 6
5 7 from ...typing import AsyncResult, Messages
6 8 from ..base_provider import AsyncGeneratorProvider, ProviderModelMixin
7 9 from ...errors import ModelNotFoundError
8 10 from ...requests import StreamSession, raise_for_status
11 from ...image import ImageResponse
9 12
10 13 from .HuggingChat import HuggingChat
11 14
12 15 class HuggingFace(AsyncGeneratorProvider, ProviderModelMixin):
13 16 url = "https://huggingface.co/chat"
14 17 working = True
15 needs_auth = True
16 18 supports_message_history = True
17 19 default_model = HuggingChat.default_model
18 models = HuggingChat.models
20 default_image_model = "black-forest-labs/FLUX.1-dev"
21 models = [*HuggingChat.models, default_image_model]
22 image_models = [default_image_model]
19 23 model_aliases = HuggingChat.model_aliases
20 24
21 25 @classmethod
@@ -29,6 +33,7 @@ class HuggingFace(AsyncGeneratorProvider, ProviderModelMixin):
29 33 api_key: str = None,
30 34 max_new_tokens: int = 1024,
31 35 temperature: float = 0.7,
36 prompt: str = None,
32 37 **kwargs
33 38 ) -> AsyncResult:
34 39 model = cls.get_model(model)
@@ -50,16 +55,22 @@ class HuggingFace(AsyncGeneratorProvider, ProviderModelMixin):
50 55 }
51 56 if api_key is not None:
52 57 headers["Authorization"] = f"Bearer {api_key}"
53 params = {
54 "return_full_text": False,
55 "max_new_tokens": max_new_tokens,
56 "temperature": temperature,
57 **kwargs
58 }
59 payload = {"inputs": format_prompt(messages), "parameters": params, "stream": stream}
58 if model in cls.image_models:
59 stream = False
60 prompt = messages[-1]["content"] if prompt is None else prompt
61 payload = {"inputs": prompt, "parameters": {"seed": random.randint(0, 2**32)}}
62 else:
63 params = {
64 "return_full_text": False,
65 "max_new_tokens": max_new_tokens,
66 "temperature": temperature,
67 **kwargs
68 }
69 payload = {"inputs": format_prompt(messages), "parameters": params, "stream": stream}
60 70 async with StreamSession(
61 71 headers=headers,
62 proxy=proxy
72 proxy=proxy,
73 timeout=600
63 74 ) as session:
64 75 async with session.post(f"{api_base.rstrip('/')}/models/{model}", json=payload) as response:
65 76 if response.status == 404:
@@ -78,7 +89,12 @@ class HuggingFace(AsyncGeneratorProvider, ProviderModelMixin):
78 89 if chunk:
79 90 yield chunk
80 91 else:
81 yield (await response.json())[0]["generated_text"].strip()
92 if response.headers["content-type"].startswith("image/"):
93 base64_data = base64.b64encode(b"".join([chunk async for chunk in response.iter_content()]))
94 url = f"data:{response.headers['content-type']};base64,{base64_data.decode()}"
95 yield ImageResponse(url, prompt)
96 else:
97 yield (await response.json())[0]["generated_text"].strip()
82 98
83 99 def format_prompt(messages: Messages) -> str:
84 100 system_messages = [message["content"] for message in messages if message["role"] == "system"]
Modified g4f/cookies.py +2 -2
@@ -34,8 +34,8 @@ try:
34 34
35 35 browsers = [
36 36 _g4f,
37 chrome, chromium, opera, opera_gx,
38 brave, edge, vivaldi, firefox,
37 chrome, chromium, firefox, opera, opera_gx,
38 brave, edge, vivaldi,
39 39 ]
40 40 has_browser_cookie3 = True
41 41 except ImportError:
Modified g4f/gui/client/static/js/chat.v1.js +3 -0
@@ -504,6 +504,8 @@ async function add_message_chunk(message, message_id) {
504 504 p.innerText = message.error;
505 505 log_storage.appendChild(p);
506 506 } else if (message.type == "preview") {
507 if (content_map.inner.clientHeight > 200)
508 content_map.inner.style.height = content_map.inner.clientHeight + "px";
507 509 content_map.inner.innerHTML = markdown_render(message.preview);
508 510 } else if (message.type == "content") {
509 511 message_storage[message_id] += message.content;
@@ -522,6 +524,7 @@ async function add_message_chunk(message, message_id) {
522 524 content_map.inner.innerHTML = html;
523 525 content_map.count.innerText = count_words_and_tokens(message_storage[message_id], provider_storage[message_id]?.model);
524 526 highlight(content_map.inner);
527 content_map.inner.style.height = "";
525 528 } else if (message.type == "log") {
526 529 let p = document.createElement("p");
527 530 p.innerText = message.log;
Modified g4f/gui/server/api.py +20 -16
@@ -123,22 +123,21 @@ class Api:
123 123 print(text)
124 124 debug.log_handler = log_handler
125 125 proxy = os.environ.get("G4F_PROXY")
126 provider = kwargs.get("provider")
127 model, provider_handler = get_model_and_provider(
128 kwargs.get("model"), provider,
129 stream=True,
130 ignore_stream=True
131 )
132 first = True
126 133 try:
127 model, provider = get_model_and_provider(
128 kwargs.get("model"), kwargs.get("provider"),
129 stream=True,
130 ignore_stream=True
131 )
132 result = ChatCompletion.create(**{**kwargs, "model": model, "provider": provider})
133 first = True
134 result = ChatCompletion.create(**{**kwargs, "model": model, "provider": provider_handler})
134 135 for chunk in result:
135 136 if first:
136 137 first = False
137 if isinstance(provider, IterListProvider):
138 provider = provider.last_provider
139 yield self._format_json("provider", {**provider.get_dict(), "model": model})
138 yield self.handle_provider(provider_handler, model)
140 139 if isinstance(chunk, BaseConversation):
141 if provider:
140 if provider is not None:
142 141 if provider not in conversations:
143 142 conversations[provider] = {}
144 143 conversations[provider][conversation_id] = chunk
@@ -165,6 +164,8 @@ class Api:
165 164 except Exception as e:
166 165 logger.exception(e)
167 166 yield self._format_json('error', get_error_message(e))
167 if first:
168 yield self.handle_provider(provider_handler, model)
168 169
169 170 def _format_json(self, response_type: str, content):
170 171 return {
@@ -172,9 +173,12 @@ class Api:
172 173 response_type: content
173 174 }
174 175
176 def handle_provider(self, provider_handler, model):
177 if isinstance(provider_handler, IterListProvider):
178 provider_handler = provider_handler.last_provider
179 if issubclass(provider_handler, ProviderModelMixin) and provider_handler.last_model is not None:
180 model = provider_handler.last_model
181 return self._format_json("provider", {**provider_handler.get_dict(), "model": model})
182
175 183 def get_error_message(exception: Exception) -> str:
176 message = f"{type(exception).__name__}: {exception}"
177 provider = get_last_provider()
178 if provider is None:
179 return message
180 return f"{provider.__name__}: {message}"
184 return f"{type(exception).__name__}: {exception}"
Modified g4f/models.py +2 -1
@@ -38,6 +38,7 @@ from .Provider import (
38 38 RubiksAI,
39 39 TeachAnything,
40 40 Upstage,
41 Flux,
41 42 )
42 43
43 44 @dataclass(unsafe_hash=True)
@@ -599,7 +600,7 @@ flux_pro = ImageModel(
599 600 flux_dev = ImageModel(
600 601 name = 'flux-dev',
601 602 base_provider = 'Flux AI',
602 best_provider = AmigoChat
603 best_provider = IterListProvider([Flux, AmigoChat, HuggingChat, HuggingFace])
603 604 )
604 605
605 606 flux_realism = ImageModel(
Modified g4f/providers/base_provider.py +3 -1
@@ -98,7 +98,7 @@ class AbstractProvider(BaseProvider):
98 98 default_value = f'"{param.default}"' if isinstance(param.default, str) else param.default
99 99 args += f" = {default_value}" if param.default is not Parameter.empty else ""
100 100 args += ","
101
101
102 102 return f"g4f.Provider.{cls.__name__} supports: ({args}\n)"
103 103
104 104 class AsyncProvider(AbstractProvider):
@@ -240,6 +240,7 @@ class ProviderModelMixin:
240 240 models: list[str] = []
241 241 model_aliases: dict[str, str] = {}
242 242 image_models: list = None
243 last_model: str = None
243 244
244 245 @classmethod
245 246 def get_models(cls) -> list[str]:
@@ -255,5 +256,6 @@ class ProviderModelMixin:
255 256 model = cls.model_aliases[model]
256 257 elif model not in cls.get_models() and cls.models:
257 258 raise ModelNotSupportedError(f"Model is not supported: {model} in: {cls.__name__}")
259 cls.last_model = model
258 260 debug.last_model = model
259 261 return model