commit
ab611018d9
@ -4,8 +4,11 @@ from src.utils import logger, get_docker_safe_url
|
|||||||
|
|
||||||
class OpenAIBase():
|
class OpenAIBase():
|
||||||
def __init__(self, api_key, base_url, model_name):
|
def __init__(self, api_key, base_url, model_name):
|
||||||
|
self.api_key = api_key
|
||||||
|
self.base_url = base_url
|
||||||
self.client = OpenAI(api_key=api_key, base_url=base_url)
|
self.client = OpenAI(api_key=api_key, base_url=base_url)
|
||||||
self.model_name = model_name
|
self.model_name = model_name
|
||||||
|
# logger.debug(f"{self.get_models()=}")
|
||||||
|
|
||||||
def predict(self, message, stream=False):
|
def predict(self, message, stream=False):
|
||||||
if isinstance(message, str):
|
if isinstance(message, str):
|
||||||
@ -34,6 +37,13 @@ class OpenAIBase():
|
|||||||
stream=False,
|
stream=False,
|
||||||
)
|
)
|
||||||
return response.choices[0].message
|
return response.choices[0].message
|
||||||
|
|
||||||
|
def get_models(self):
|
||||||
|
try:
|
||||||
|
return self.client.models.list()
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Error getting models: {e}")
|
||||||
|
return []
|
||||||
|
|
||||||
|
|
||||||
class OpenModel(OpenAIBase):
|
class OpenModel(OpenAIBase):
|
||||||
@ -60,7 +70,8 @@ class GeneralResponse:
|
|||||||
self.is_full = False
|
self.is_full = False
|
||||||
|
|
||||||
|
|
||||||
class Qianfan:
|
class Qianfan(OpenAIBase):
|
||||||
|
"""弃用"""
|
||||||
|
|
||||||
def __init__(self, model_name="ernie_speed") -> None:
|
def __init__(self, model_name="ernie_speed") -> None:
|
||||||
import qianfan
|
import qianfan
|
||||||
@ -99,7 +110,7 @@ class Qianfan:
|
|||||||
|
|
||||||
|
|
||||||
|
|
||||||
class DashScope:
|
class DashScope(OpenAIBase):
|
||||||
|
|
||||||
def __init__(self, model_name="qwen-max-latest") -> None:
|
def __init__(self, model_name="qwen-max-latest") -> None:
|
||||||
self.model_name = model_name
|
self.model_name = model_name
|
||||||
@ -143,6 +154,4 @@ class DashScope:
|
|||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
model = SiliconFlow()
|
pass
|
||||||
for a in model.predict("你好", stream=True):
|
|
||||||
print(a.content, end="")
|
|
||||||
@ -21,6 +21,7 @@ def chat_post(
|
|||||||
history: list = Body(...),
|
history: list = Body(...),
|
||||||
cur_res_id: str = Body(...)):
|
cur_res_id: str = Body(...)):
|
||||||
|
|
||||||
|
meta["server_model_name"] = startup.model.model_name
|
||||||
history_manager = HistoryManager(history)
|
history_manager = HistoryManager(history)
|
||||||
logger.debug(f"Received query: {query} with meta: {meta}")
|
logger.debug(f"Received query: {query} with meta: {meta}")
|
||||||
|
|
||||||
|
|||||||
@ -3,7 +3,7 @@
|
|||||||
<div class="tags">
|
<div class="tags">
|
||||||
<!-- <span class="item btn" @click="likeThisResponse(msg)"><LikeOutlined /></span> -->
|
<!-- <span class="item btn" @click="likeThisResponse(msg)"><LikeOutlined /></span> -->
|
||||||
<!-- <span class="item btn" @click="dislikeThisResponse(msg)"><DislikeOutlined /></span> -->
|
<!-- <span class="item btn" @click="dislikeThisResponse(msg)"><DislikeOutlined /></span> -->
|
||||||
<span class="item"><BulbOutlined /> {{ msg.model_name }}</span>
|
<span class="item"><BulbOutlined /> {{ msg.meta.server_model_name }}</span>
|
||||||
<span class="item btn" @click="copyText(msg.text)"><CopyOutlined /></span>
|
<span class="item btn" @click="copyText(msg.text)"><CopyOutlined /></span>
|
||||||
<span
|
<span
|
||||||
class="item btn"
|
class="item btn"
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user