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

View File

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

View File

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

View File

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