185 lines
7.1 KiB
Python
185 lines
7.1 KiB
Python
import os
|
||
import pickle
|
||
from pathlib import Path
|
||
from typing import List
|
||
|
||
import faiss
|
||
import numpy as np
|
||
import streamlit as st
|
||
from langchain_core.embeddings import Embeddings
|
||
from openai import OpenAI
|
||
from sentence_transformers import SentenceTransformer, CrossEncoder
|
||
|
||
LM_STUDIO_URL = os.getenv("LM_STUDIO_URL", "http://localhost:1234")
|
||
EMBED_MODEL = "intfloat/multilingual-e5-large"
|
||
VECTOR_DIR = Path(__file__).parent / "vector_store"
|
||
VECTOR_INDEX = VECTOR_DIR / "index.faiss.bin"
|
||
VECTOR_META = VECTOR_DIR / "metadata.pkl"
|
||
|
||
SYSTEM_PROMPT = """Ты — специалист техподдержки. Отвечай клиенту, используя ТОЛЬКО информацию из переданных тикетов (контекст).
|
||
|
||
Правила:
|
||
- Если контекст содержит подходящее решение — напиши ответ своими словами, адаптируя под вопрос
|
||
- Если контекст не относится к вопросу — напиши: «Недостаточно информации в истории обращений»
|
||
- НЕ придумывай ответы, НЕ используй общие знания
|
||
- НЕ говори «обратитесь в службу поддержки» — ты сам и есть поддержка
|
||
- Укажи в конце: «Основано на тикете №...»"""
|
||
|
||
FETCH_K = 20 # сколько достаём из FAISS для реранжа
|
||
FINAL_K = 5 # сколько оставляем после реранжа
|
||
|
||
|
||
class LocalEmbeddings(Embeddings):
|
||
def __init__(self, model_name: str):
|
||
self.model = SentenceTransformer(model_name, device="cuda")
|
||
|
||
def embed_documents(self, texts: List[str]) -> List[List[float]]:
|
||
emb = self.model.encode(texts, normalize_embeddings=True)
|
||
return emb.tolist()
|
||
|
||
def embed_query(self, text: str) -> List[float]:
|
||
emb = self.model.encode([text], normalize_embeddings=True)
|
||
return emb[0].tolist()
|
||
|
||
|
||
@st.cache_resource
|
||
def load_embeddings():
|
||
return LocalEmbeddings(EMBED_MODEL)
|
||
|
||
|
||
@st.cache_resource
|
||
def load_reranker():
|
||
return CrossEncoder(
|
||
"cross-encoder/mmarco-mMiniLMv2-L12-H384-v1",
|
||
max_length=512,
|
||
device="cpu",
|
||
)
|
||
|
||
|
||
def load_faiss_index(path: Path):
|
||
return faiss.deserialize_index(np.frombuffer(path.read_bytes(), dtype=np.uint8))
|
||
|
||
|
||
def get_chat_client() -> OpenAI | None:
|
||
try:
|
||
OpenAI(base_url=f"{LM_STUDIO_URL}/v1", api_key="not-needed").models.list()
|
||
return OpenAI(base_url=f"{LM_STUDIO_URL}/v1", api_key="not-needed")
|
||
except Exception:
|
||
return None
|
||
|
||
|
||
def search_similar(query: str, k: int = FETCH_K):
|
||
embeddings = load_embeddings()
|
||
qvec = np.array([embeddings.embed_query(query)], dtype=np.float32)
|
||
index = load_faiss_index(VECTOR_INDEX)
|
||
scores, indices = index.search(qvec, k)
|
||
with open(VECTOR_META, "rb") as f:
|
||
data = pickle.load(f)
|
||
metadatas = data["metadatas"]
|
||
results = []
|
||
for score, idx in zip(scores[0], indices[0]):
|
||
if idx < 0 or idx >= len(metadatas):
|
||
continue
|
||
results.append((metadatas[idx], float(score)))
|
||
results.sort(key=lambda x: x[1], reverse=True)
|
||
return results
|
||
|
||
|
||
def rerank(query: str, results: list) -> list:
|
||
reranker = load_reranker()
|
||
pairs = [(query, r[0]["full_text"]) for r in results]
|
||
scores = reranker.predict(pairs)
|
||
scored = [(results[i][0], float(scores[i])) for i in range(len(results))]
|
||
scored.sort(key=lambda x: x[1], reverse=True)
|
||
return scored[:FINAL_K]
|
||
|
||
|
||
def format_context(results) -> str:
|
||
parts = []
|
||
for meta, score in results:
|
||
parts.append(
|
||
f"Тикет №{meta['ticket_id']} ({meta['category']}, {meta['client']}):\n{meta['full_text']}"
|
||
)
|
||
return "\n\n---\n\n".join(parts)
|
||
|
||
|
||
def generate_answer(chat: OpenAI, query: str, context: str) -> str:
|
||
response = chat.chat.completions.create(
|
||
model="gpt-3.5-turbo",
|
||
messages=[
|
||
{"role": "system", "content": SYSTEM_PROMPT},
|
||
{"role": "user", "content": f"Контекст (история обращений):\n{context}\n\nВопрос клиента:\n{query}"},
|
||
],
|
||
temperature=0.1,
|
||
max_tokens=2000,
|
||
)
|
||
return response.choices[0].message.content
|
||
|
||
|
||
def main():
|
||
st.set_page_config(page_title="RAG техподдержки", layout="wide")
|
||
st.title("RAG-ассистент техподдержки")
|
||
st.markdown(
|
||
"Находит похожие обращения из истории и формирует ответ на их основе."
|
||
)
|
||
|
||
if not VECTOR_INDEX.exists():
|
||
st.error("Векторный индекс не найден. Запусти: python prepare_data.py")
|
||
st.stop()
|
||
|
||
chat = get_chat_client()
|
||
if chat is None:
|
||
st.warning(
|
||
"⚠️ Чат-модель не обнаружена. Загрузи модель в LM Studio."
|
||
)
|
||
|
||
with st.form("query_form"):
|
||
query = st.text_area(
|
||
"Опишите проблему клиента:",
|
||
placeholder="Например: клиент не может войти в ТГ-бот, пишет неверный логин",
|
||
height=100,
|
||
)
|
||
submitted = st.form_submit_button("Получить ответ", type="primary")
|
||
|
||
if submitted and query:
|
||
with st.spinner("Ищу похожие обращения..."):
|
||
results = search_similar(query)
|
||
|
||
if not results:
|
||
st.warning("Не найдено похожих обращений")
|
||
return
|
||
|
||
with st.spinner("Оцениваю релевантность..."):
|
||
results = rerank(query, results)
|
||
|
||
if not results:
|
||
st.warning("После проверки релевантности не осталось подходящих тикетов")
|
||
return
|
||
|
||
context = format_context(results)
|
||
|
||
if chat:
|
||
with st.spinner("Формирую ответ на основе найденных тикетов..."):
|
||
try:
|
||
answer = generate_answer(chat, query, context)
|
||
st.subheader("💬 Ответ ассистента")
|
||
st.markdown(answer)
|
||
except Exception as e:
|
||
st.error(f"Ошибка LLM: {e}")
|
||
else:
|
||
st.info("Результаты поиска без генерации ответа:")
|
||
|
||
with st.expander("📄 Найденные похожие обращения", expanded=True):
|
||
for meta, score in results:
|
||
with st.container(border=True):
|
||
st.markdown(f"### Тикет №{meta['ticket_id']} (релевантность: {score:.3f})")
|
||
col1, col2 = st.columns(2)
|
||
with col1:
|
||
st.markdown(f"**Категория:** {meta['category']}")
|
||
with col2:
|
||
st.markdown(f"**Клиент:** {meta['client']}")
|
||
st.text(meta["full_text"])
|
||
|
||
|
||
if __name__ == "__main__":
|
||
main()
|