Add support of temperature and max tokens parameters

This commit is contained in:
dimitrievgs 2026-05-11 16:57:14 +03:00
parent 526f4e7b94
commit 42469c28d1
4 changed files with 63 additions and 16 deletions

View File

@ -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}")

View File

@ -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

View File

@ -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 для рендеринга на фронте

View File

@ -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}