Add support of temperature and max tokens parameters
This commit is contained in:
parent
526f4e7b94
commit
42469c28d1
12
app/api.py
12
app/api.py
|
|
@ -9,7 +9,7 @@ from flask import Flask, request, jsonify, Response, stream_with_context
|
||||||
from flask_cors import CORS
|
from flask_cors import CORS
|
||||||
from flask_socketio import SocketIO
|
from flask_socketio import SocketIO
|
||||||
from workflows import graph_history_manager, run_agent_streaming
|
from workflows import graph_history_manager, run_agent_streaming
|
||||||
from llm_client import MODELS, get_llm # Добавляем импорт списка моделей
|
from llm_client import DEFAULT_TEMPERATURE, MODELS, get_llm # Добавляем импорт списка моделей
|
||||||
|
|
||||||
api = Flask(__name__)
|
api = Flask(__name__)
|
||||||
CORS(
|
CORS(
|
||||||
|
|
@ -229,6 +229,8 @@ def chat_stream():
|
||||||
system_prompt = data.get("system_prompt")
|
system_prompt = data.get("system_prompt")
|
||||||
model = data.get("model")
|
model = data.get("model")
|
||||||
cache_folder = data.get("cache_folder")
|
cache_folder = data.get("cache_folder")
|
||||||
|
temperature = data.get("temperature", DEFAULT_TEMPERATURE)
|
||||||
|
max_tokens = data.get("max_tokens")
|
||||||
# Необязательный параметр для существующего узла ассистента при регенерации
|
# Необязательный параметр для существующего узла ассистента при регенерации
|
||||||
existing_assistant_node_id = data.get("existing_assistant_node_id")
|
existing_assistant_node_id = data.get("existing_assistant_node_id")
|
||||||
|
|
||||||
|
|
@ -257,7 +259,7 @@ def chat_stream():
|
||||||
for chunk_data in run_agent_streaming(graph_id, user_node_id,
|
for chunk_data in run_agent_streaming(graph_id, user_node_id,
|
||||||
assistant_node_id,
|
assistant_node_id,
|
||||||
system_prompt, model,
|
system_prompt, model,
|
||||||
cache_folder):
|
cache_folder, temperature, max_tokens):
|
||||||
if chunk_data.get("type") == "chunk":
|
if chunk_data.get("type") == "chunk":
|
||||||
accumulated_content += chunk_data.get("content", "")
|
accumulated_content += chunk_data.get("content", "")
|
||||||
yield f"data: {json.dumps(chunk_data)}\n\n"
|
yield f"data: {json.dumps(chunk_data)}\n\n"
|
||||||
|
|
@ -328,6 +330,8 @@ def regenerate_message():
|
||||||
model = data.get("model")
|
model = data.get("model")
|
||||||
system_prompt = data.get("system_prompt")
|
system_prompt = data.get("system_prompt")
|
||||||
cache_folder = data.get("cache_folder")
|
cache_folder = data.get("cache_folder")
|
||||||
|
temperature = data.get("temperature", DEFAULT_TEMPERATURE)
|
||||||
|
max_tokens = data.get("max_tokens")
|
||||||
|
|
||||||
if not graph_id or not node_id:
|
if not graph_id or not node_id:
|
||||||
return jsonify({"error": "Не указаны graph_id или node_id."}, 400)
|
return jsonify({"error": "Не указаны graph_id или node_id."}, 400)
|
||||||
|
|
@ -362,7 +366,9 @@ def regenerate_message():
|
||||||
"graph_id": graph_id,
|
"graph_id": graph_id,
|
||||||
"model": model,
|
"model": model,
|
||||||
"system_prompt": system_prompt,
|
"system_prompt": system_prompt,
|
||||||
"cache_folder": cache_folder
|
"cache_folder": cache_folder,
|
||||||
|
"temperature": temperature,
|
||||||
|
"max_tokens": max_tokens
|
||||||
})
|
})
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
print(f"Ошибка при регенерации сообщения: {e}")
|
print(f"Ошибка при регенерации сообщения: {e}")
|
||||||
|
|
|
||||||
|
|
@ -3,7 +3,7 @@
|
||||||
включая OpenAI (Gemini) и MistralAI, обеспечивая унифицированный интерфейс.
|
включая OpenAI (Gemini) и MistralAI, обеспечивая унифицированный интерфейс.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from typing import Dict, Any, List
|
from typing import Dict, Any, List, Optional
|
||||||
from langchain_core.messages import HumanMessage, SystemMessage, AIMessage
|
from langchain_core.messages import HumanMessage, SystemMessage, AIMessage
|
||||||
|
|
||||||
from openai import OpenAI
|
from openai import OpenAI
|
||||||
|
|
@ -14,6 +14,8 @@ from mistralai import Mistral
|
||||||
# Пустой словарь для локальных моделей (Ollama)
|
# Пустой словарь для локальных моделей (Ollama)
|
||||||
LOCAL_MODELS: Dict[str, Any] = {}
|
LOCAL_MODELS: Dict[str, Any] = {}
|
||||||
|
|
||||||
|
DEFAULT_TEMPERATURE = 1.0
|
||||||
|
|
||||||
# Конфигурация внешних моделей
|
# Конфигурация внешних моделей
|
||||||
MODELS: Dict[str, Dict[str, Any]] = {
|
MODELS: Dict[str, Dict[str, Any]] = {
|
||||||
"gemini-2.0-flash": {
|
"gemini-2.0-flash": {
|
||||||
|
|
@ -215,7 +217,7 @@ class CustomLLM:
|
||||||
print(f"Ошибка при вызове OpenAI API: {e}")
|
print(f"Ошибка при вызове OpenAI API: {e}")
|
||||||
raise
|
raise
|
||||||
|
|
||||||
def invoke(self, messages: List[Any]) -> str:
|
def invoke(self, messages: List[Any], temperature: float = DEFAULT_TEMPERATURE, max_tokens: Optional[int] = None) -> str:
|
||||||
"""
|
"""
|
||||||
Отправляет запрос к LLM и возвращает ответ.
|
Отправляет запрос к LLM и возвращает ответ.
|
||||||
Messages могут быть строкой или списком объектов Langchain Message.
|
Messages могут быть строкой или списком объектов Langchain Message.
|
||||||
|
|
@ -252,10 +254,15 @@ class CustomLLM:
|
||||||
})
|
})
|
||||||
|
|
||||||
try:
|
try:
|
||||||
|
kwargs = {"temperature": temperature}
|
||||||
|
if max_tokens:
|
||||||
|
kwargs["max_tokens"] = int(max_tokens)
|
||||||
|
|
||||||
response = self.client.chat.completions.create(
|
response = self.client.chat.completions.create(
|
||||||
model=self.model_name,
|
model=self.model_name,
|
||||||
messages=openai_messages,
|
messages=openai_messages,
|
||||||
stream=self.stream,
|
stream=self.stream,
|
||||||
|
**kwargs
|
||||||
)
|
)
|
||||||
|
|
||||||
if self.stream:
|
if self.stream:
|
||||||
|
|
@ -309,8 +316,13 @@ class CustomLLM:
|
||||||
})
|
})
|
||||||
|
|
||||||
try:
|
try:
|
||||||
|
kwargs = {"temperature": temperature}
|
||||||
|
if max_tokens:
|
||||||
|
kwargs["max_tokens"] = int(max_tokens)
|
||||||
|
|
||||||
response = self.client.chat.stream(model=self.model_name,
|
response = self.client.chat.stream(model=self.model_name,
|
||||||
messages=mistralai_messages)
|
messages=mistralai_messages,
|
||||||
|
**kwargs)
|
||||||
|
|
||||||
if self.stream:
|
if self.stream:
|
||||||
collected_messages = []
|
collected_messages = []
|
||||||
|
|
@ -336,7 +348,7 @@ class CustomLLM:
|
||||||
raise
|
raise
|
||||||
return ""
|
return ""
|
||||||
|
|
||||||
def stream2(self, messages: List[Any]):
|
def stream2(self, messages: List[Any], temperature: float = DEFAULT_TEMPERATURE, max_tokens: Optional[int] = None):
|
||||||
"""
|
"""
|
||||||
Генератор для стриминга ответа от LLM.
|
Генератор для стриминга ответа от LLM.
|
||||||
Возвращает чанки контента по мере их получения.
|
Возвращает чанки контента по мере их получения.
|
||||||
|
|
@ -345,10 +357,15 @@ class CustomLLM:
|
||||||
openai_messages = self._prepare_openai_messages(messages)
|
openai_messages = self._prepare_openai_messages(messages)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
|
kwargs = {"temperature": temperature}
|
||||||
|
if max_tokens:
|
||||||
|
kwargs["max_tokens"] = int(max_tokens)
|
||||||
|
|
||||||
response = self.client.chat.completions.create(
|
response = self.client.chat.completions.create(
|
||||||
model=self.model_name,
|
model=self.model_name,
|
||||||
messages=openai_messages,
|
messages=openai_messages,
|
||||||
stream=True,
|
stream=True,
|
||||||
|
**kwargs
|
||||||
)
|
)
|
||||||
|
|
||||||
for chunk in response:
|
for chunk in response:
|
||||||
|
|
@ -368,8 +385,13 @@ class CustomLLM:
|
||||||
mistralai_messages = self._prepare_mistralai_messages(messages)
|
mistralai_messages = self._prepare_mistralai_messages(messages)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
|
kwargs = {"temperature": temperature}
|
||||||
|
if max_tokens:
|
||||||
|
kwargs["max_tokens"] = int(max_tokens)
|
||||||
|
|
||||||
response = self.client.chat.stream(model=self.model_name,
|
response = self.client.chat.stream(model=self.model_name,
|
||||||
messages=mistralai_messages)
|
messages=mistralai_messages,
|
||||||
|
**kwargs)
|
||||||
|
|
||||||
for chunk in response:
|
for chunk in response:
|
||||||
delta = chunk.data.choices[0].delta
|
delta = chunk.data.choices[0].delta
|
||||||
|
|
|
||||||
|
|
@ -7,7 +7,7 @@ import sqlite3
|
||||||
import threading
|
import threading
|
||||||
import time
|
import time
|
||||||
from typing import Any, Dict, Optional
|
from typing import Any, Dict, Optional
|
||||||
from llm_client import get_llm
|
from llm_client import DEFAULT_TEMPERATURE, get_llm
|
||||||
from langchain_core.messages import SystemMessage, HumanMessage
|
from langchain_core.messages import SystemMessage, HumanMessage
|
||||||
import uuid
|
import uuid
|
||||||
|
|
||||||
|
|
@ -17,6 +17,9 @@ DEFAULT_VOICE_LLM_NAME = "gemini-2.0-flash-lite" # "gemini-3.0-flash-openrouter"
|
||||||
MAX_TITLE_GENERATION_CONTENT_LENGTH = 5000 # Максимальное количество символов для генерации заголовков
|
MAX_TITLE_GENERATION_CONTENT_LENGTH = 5000 # Максимальное количество символов для генерации заголовков
|
||||||
QUEUE_POLLING_INTERVAL = 4
|
QUEUE_POLLING_INTERVAL = 4
|
||||||
|
|
||||||
|
TITLE_TEMPERATURE = 0.1
|
||||||
|
MAX_TITLE_LENGTH = 100
|
||||||
|
|
||||||
from app.voice_service import VOICE_COMMANDS_RESPONSE_TO_STORE
|
from app.voice_service import VOICE_COMMANDS_RESPONSE_TO_STORE
|
||||||
|
|
||||||
class TitleGenerator:
|
class TitleGenerator:
|
||||||
|
|
@ -203,7 +206,7 @@ class TitleGenerator:
|
||||||
HumanMessage(content=first_user_message)
|
HumanMessage(content=first_user_message)
|
||||||
]
|
]
|
||||||
|
|
||||||
title = self.llm.invoke(llm_messages)
|
title = self.llm.invoke(llm_messages, temperature=TITLE_TEMPERATURE, max_tokens=MAX_TITLE_LENGTH)
|
||||||
title = title.strip()[:80] # Ограничиваем длину
|
title = title.strip()[:80] # Ограничиваем длину
|
||||||
|
|
||||||
# Сохраняем заголовок
|
# Сохраняем заголовок
|
||||||
|
|
@ -257,7 +260,7 @@ class TitleGenerator:
|
||||||
HumanMessage(content=content)
|
HumanMessage(content=content)
|
||||||
]
|
]
|
||||||
|
|
||||||
title = self.llm.invoke(llm_messages)
|
title = self.llm.invoke(llm_messages, temperature=TITLE_TEMPERATURE, max_tokens=MAX_TITLE_LENGTH)
|
||||||
title = title.strip()[:40] # Ограничиваем длину
|
title = title.strip()[:40] # Ограничиваем длину
|
||||||
|
|
||||||
# Сохраняем заголовок
|
# Сохраняем заголовок
|
||||||
|
|
@ -379,12 +382,19 @@ class TitleGenerator:
|
||||||
# 3. Добавьте метод обработки регулярной задачи
|
# 3. Добавьте метод обработки регулярной задачи
|
||||||
def _process_voice_regular(self, text, prompt):
|
def _process_voice_regular(self, text, prompt):
|
||||||
try:
|
try:
|
||||||
|
import app.api as api
|
||||||
|
settings = api.current_obsidian_settings
|
||||||
|
temp = float(settings.get('temperatureVoice', DEFAULT_TEMPERATURE))
|
||||||
|
max_t = settings.get('maxTokensVoiceRegular')
|
||||||
|
if max_t == "" or max_t == 0: max_t = None
|
||||||
|
elif max_t is not None: max_t = int(max_t)
|
||||||
|
|
||||||
print(f"🎤 Анализ транскрибации (Prompt 1)...")
|
print(f"🎤 Анализ транскрибации (Prompt 1)...")
|
||||||
llm_messages = [
|
llm_messages = [
|
||||||
SystemMessage(content=prompt),
|
SystemMessage(content=prompt),
|
||||||
HumanMessage(content=f"Последний фрагмент диалога для анализа: {text}")
|
HumanMessage(content=f"Последний фрагмент диалога для анализа: {text}")
|
||||||
]
|
]
|
||||||
response = self.voice_llm.invoke(llm_messages)
|
response = self.voice_llm.invoke(llm_messages, temperature=temp, max_tokens=max_t)
|
||||||
|
|
||||||
import app.api as api
|
import app.api as api
|
||||||
new_entry = {
|
new_entry = {
|
||||||
|
|
@ -401,12 +411,19 @@ class TitleGenerator:
|
||||||
# 4. Поправьте _process_voice_command (запись команды в Response 2)
|
# 4. Поправьте _process_voice_command (запись команды в Response 2)
|
||||||
def _process_voice_command(self, text, prompt):
|
def _process_voice_command(self, text, prompt):
|
||||||
try:
|
try:
|
||||||
|
import app.api as api
|
||||||
|
settings = api.current_obsidian_settings
|
||||||
|
temp = float(settings.get('temperatureVoice', DEFAULT_TEMPERATURE))
|
||||||
|
max_t = settings.get('maxTokensVoiceCommands')
|
||||||
|
if max_t == "" or max_t == 0: max_t = None
|
||||||
|
elif max_t is not None: max_t = int(max_t)
|
||||||
|
|
||||||
print(f"🎤 Обработка голосовой команды (Prompt 2)...")
|
print(f"🎤 Обработка голосовой команды (Prompt 2)...")
|
||||||
llm_messages = [
|
llm_messages = [
|
||||||
SystemMessage(content=prompt),
|
SystemMessage(content=prompt),
|
||||||
HumanMessage(content=f"Выполни команду из транскрибации: {text}")
|
HumanMessage(content=f"Выполни команду из транскрибации: {text}")
|
||||||
]
|
]
|
||||||
response = self.voice_llm.invoke(llm_messages)
|
response = self.voice_llm.invoke(llm_messages, temperature=temp, max_tokens=max_t)
|
||||||
|
|
||||||
import app.api as api
|
import app.api as api
|
||||||
# Форматируем как Markdown для рендеринга на фронте
|
# Форматируем как Markdown для рендеринга на фронте
|
||||||
|
|
|
||||||
|
|
@ -10,7 +10,7 @@ from graph_history_manager import GraphHistoryManager
|
||||||
from models import AgentState
|
from models import AgentState
|
||||||
from nodes import parse_command_node, execute_command_node, call_llm_node, generate_images_node, analyze_image_node, get_meet_subtitles_node, get_teams_subtitles_node, summarize_history_node, handle_error_node, help_node
|
from nodes import parse_command_node, execute_command_node, call_llm_node, generate_images_node, analyze_image_node, get_meet_subtitles_node, get_teams_subtitles_node, summarize_history_node, handle_error_node, help_node
|
||||||
import uuid # Добавлено для генерации UUID
|
import uuid # Добавлено для генерации UUID
|
||||||
from llm_client import get_llm
|
from llm_client import DEFAULT_TEMPERATURE, get_llm
|
||||||
|
|
||||||
import os
|
import os
|
||||||
import base64
|
import base64
|
||||||
|
|
@ -100,7 +100,9 @@ def run_agent_streaming(graph_id: str,
|
||||||
assistant_node_id: str,
|
assistant_node_id: str,
|
||||||
system_prompt: Optional[str] = None,
|
system_prompt: Optional[str] = None,
|
||||||
model: Optional[str] = None,
|
model: Optional[str] = None,
|
||||||
cache_folder: Optional[str] = None):
|
cache_folder: Optional[str] = None,
|
||||||
|
temperature: float = DEFAULT_TEMPERATURE,
|
||||||
|
max_tokens: Optional[int] = None):
|
||||||
"""Генератор для стримингового ответа LLM с поддержкой вложений."""
|
"""Генератор для стримингового ответа LLM с поддержкой вложений."""
|
||||||
try:
|
try:
|
||||||
loaded_graph_data = graph_history_manager.get_graph(
|
loaded_graph_data = graph_history_manager.get_graph(
|
||||||
|
|
@ -163,7 +165,7 @@ def run_agent_streaming(graph_id: str,
|
||||||
|
|
||||||
# Стриминг ответа
|
# Стриминг ответа
|
||||||
llm = get_llm(model or "gemini-2.5-flash")
|
llm = get_llm(model or "gemini-2.5-flash")
|
||||||
for chunk in llm.stream2(messages_for_llm):
|
for chunk in llm.stream2(messages_for_llm, temperature=temperature, max_tokens=max_tokens):
|
||||||
content = chunk.content if hasattr(chunk,
|
content = chunk.content if hasattr(chunk,
|
||||||
'content') else str(chunk)
|
'content') else str(chunk)
|
||||||
yield {"type": "chunk", "content": content}
|
yield {"type": "chunk", "content": content}
|
||||||
|
|
|
||||||
Loading…
Reference in New Issue
Block a user