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_socketio import SocketIO
|
||||
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__)
|
||||
CORS(
|
||||
|
|
@ -229,6 +229,8 @@ def chat_stream():
|
|||
system_prompt = data.get("system_prompt")
|
||||
model = data.get("model")
|
||||
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")
|
||||
|
||||
|
|
@ -257,7 +259,7 @@ def chat_stream():
|
|||
for chunk_data in run_agent_streaming(graph_id, user_node_id,
|
||||
assistant_node_id,
|
||||
system_prompt, model,
|
||||
cache_folder):
|
||||
cache_folder, temperature, max_tokens):
|
||||
if chunk_data.get("type") == "chunk":
|
||||
accumulated_content += chunk_data.get("content", "")
|
||||
yield f"data: {json.dumps(chunk_data)}\n\n"
|
||||
|
|
@ -328,6 +330,8 @@ def regenerate_message():
|
|||
model = data.get("model")
|
||||
system_prompt = data.get("system_prompt")
|
||||
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:
|
||||
return jsonify({"error": "Не указаны graph_id или node_id."}, 400)
|
||||
|
|
@ -362,7 +366,9 @@ def regenerate_message():
|
|||
"graph_id": graph_id,
|
||||
"model": model,
|
||||
"system_prompt": system_prompt,
|
||||
"cache_folder": cache_folder
|
||||
"cache_folder": cache_folder,
|
||||
"temperature": temperature,
|
||||
"max_tokens": max_tokens
|
||||
})
|
||||
except Exception as e:
|
||||
print(f"Ошибка при регенерации сообщения: {e}")
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@
|
|||
включая 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 openai import OpenAI
|
||||
|
|
@ -14,6 +14,8 @@ from mistralai import Mistral
|
|||
# Пустой словарь для локальных моделей (Ollama)
|
||||
LOCAL_MODELS: Dict[str, Any] = {}
|
||||
|
||||
DEFAULT_TEMPERATURE = 1.0
|
||||
|
||||
# Конфигурация внешних моделей
|
||||
MODELS: Dict[str, Dict[str, Any]] = {
|
||||
"gemini-2.0-flash": {
|
||||
|
|
@ -215,7 +217,7 @@ class CustomLLM:
|
|||
print(f"Ошибка при вызове OpenAI API: {e}")
|
||||
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 и возвращает ответ.
|
||||
Messages могут быть строкой или списком объектов Langchain Message.
|
||||
|
|
@ -252,10 +254,15 @@ class CustomLLM:
|
|||
})
|
||||
|
||||
try:
|
||||
kwargs = {"temperature": temperature}
|
||||
if max_tokens:
|
||||
kwargs["max_tokens"] = int(max_tokens)
|
||||
|
||||
response = self.client.chat.completions.create(
|
||||
model=self.model_name,
|
||||
messages=openai_messages,
|
||||
stream=self.stream,
|
||||
**kwargs
|
||||
)
|
||||
|
||||
if self.stream:
|
||||
|
|
@ -309,8 +316,13 @@ class CustomLLM:
|
|||
})
|
||||
|
||||
try:
|
||||
kwargs = {"temperature": temperature}
|
||||
if max_tokens:
|
||||
kwargs["max_tokens"] = int(max_tokens)
|
||||
|
||||
response = self.client.chat.stream(model=self.model_name,
|
||||
messages=mistralai_messages)
|
||||
messages=mistralai_messages,
|
||||
**kwargs)
|
||||
|
||||
if self.stream:
|
||||
collected_messages = []
|
||||
|
|
@ -336,7 +348,7 @@ class CustomLLM:
|
|||
raise
|
||||
return ""
|
||||
|
||||
def stream2(self, messages: List[Any]):
|
||||
def stream2(self, messages: List[Any], temperature: float = DEFAULT_TEMPERATURE, max_tokens: Optional[int] = None):
|
||||
"""
|
||||
Генератор для стриминга ответа от LLM.
|
||||
Возвращает чанки контента по мере их получения.
|
||||
|
|
@ -345,10 +357,15 @@ class CustomLLM:
|
|||
openai_messages = self._prepare_openai_messages(messages)
|
||||
|
||||
try:
|
||||
kwargs = {"temperature": temperature}
|
||||
if max_tokens:
|
||||
kwargs["max_tokens"] = int(max_tokens)
|
||||
|
||||
response = self.client.chat.completions.create(
|
||||
model=self.model_name,
|
||||
messages=openai_messages,
|
||||
stream=True,
|
||||
**kwargs
|
||||
)
|
||||
|
||||
for chunk in response:
|
||||
|
|
@ -368,8 +385,13 @@ class CustomLLM:
|
|||
mistralai_messages = self._prepare_mistralai_messages(messages)
|
||||
|
||||
try:
|
||||
kwargs = {"temperature": temperature}
|
||||
if max_tokens:
|
||||
kwargs["max_tokens"] = int(max_tokens)
|
||||
|
||||
response = self.client.chat.stream(model=self.model_name,
|
||||
messages=mistralai_messages)
|
||||
messages=mistralai_messages,
|
||||
**kwargs)
|
||||
|
||||
for chunk in response:
|
||||
delta = chunk.data.choices[0].delta
|
||||
|
|
|
|||
|
|
@ -7,7 +7,7 @@ import sqlite3
|
|||
import threading
|
||||
import time
|
||||
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
|
||||
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 # Максимальное количество символов для генерации заголовков
|
||||
QUEUE_POLLING_INTERVAL = 4
|
||||
|
||||
TITLE_TEMPERATURE = 0.1
|
||||
MAX_TITLE_LENGTH = 100
|
||||
|
||||
from app.voice_service import VOICE_COMMANDS_RESPONSE_TO_STORE
|
||||
|
||||
class TitleGenerator:
|
||||
|
|
@ -203,7 +206,7 @@ class TitleGenerator:
|
|||
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] # Ограничиваем длину
|
||||
|
||||
# Сохраняем заголовок
|
||||
|
|
@ -257,7 +260,7 @@ class TitleGenerator:
|
|||
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] # Ограничиваем длину
|
||||
|
||||
# Сохраняем заголовок
|
||||
|
|
@ -379,12 +382,19 @@ class TitleGenerator:
|
|||
# 3. Добавьте метод обработки регулярной задачи
|
||||
def _process_voice_regular(self, text, prompt):
|
||||
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)...")
|
||||
llm_messages = [
|
||||
SystemMessage(content=prompt),
|
||||
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
|
||||
new_entry = {
|
||||
|
|
@ -401,12 +411,19 @@ class TitleGenerator:
|
|||
# 4. Поправьте _process_voice_command (запись команды в Response 2)
|
||||
def _process_voice_command(self, text, prompt):
|
||||
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)...")
|
||||
llm_messages = [
|
||||
SystemMessage(content=prompt),
|
||||
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
|
||||
# Форматируем как Markdown для рендеринга на фронте
|
||||
|
|
|
|||
|
|
@ -10,7 +10,7 @@ from graph_history_manager import GraphHistoryManager
|
|||
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
|
||||
import uuid # Добавлено для генерации UUID
|
||||
from llm_client import get_llm
|
||||
from llm_client import DEFAULT_TEMPERATURE, get_llm
|
||||
|
||||
import os
|
||||
import base64
|
||||
|
|
@ -100,7 +100,9 @@ def run_agent_streaming(graph_id: str,
|
|||
assistant_node_id: str,
|
||||
system_prompt: 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 с поддержкой вложений."""
|
||||
try:
|
||||
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")
|
||||
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') else str(chunk)
|
||||
yield {"type": "chunk", "content": content}
|
||||
|
|
|
|||
Loading…
Reference in New Issue
Block a user