X commited on
Commit
b65f3c6
·
verified ·
1 Parent(s): 61e218f

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +49 -60
app.py CHANGED
@@ -24,7 +24,8 @@ print(f"CUDA доступна: {torch.cuda.is_available()}")
24
  if torch.cuda.is_available():
25
  print(f"GPU: {torch.cuda.get_device_name(0)}")
26
 
27
- MODEL_NAME = "ai-forever/rugpt3small_based_on_gpt2"
 
28
  EMBEDDING_MODEL = "all-MiniLM-L6-v2"
29
  SCIENCE_DATASET = "RafaelUI/ru_science"
30
  ARTICLE_LIMIT = 50
@@ -51,7 +52,19 @@ st.set_page_config(
51
  )
52
 
53
  # ===================================================================
54
- # 2. НАСТОЯЩАЯ НЕЙРОСЕТЬ (ВСЕГДА ГЕНЕРИРУЕТ)
 
 
 
 
 
 
 
 
 
 
 
 
55
  # ===================================================================
56
 
57
  class NeuralChatbot:
@@ -62,7 +75,6 @@ class NeuralChatbot:
62
  self.generator = None
63
  self.is_loaded = False
64
 
65
- # Системный промпт для нейросети
66
  self.system_prompt = f"""Ты - {AI_NAME}, дружелюбный научный AI-ассистент от компании {COMPANY_NAME}.
67
  Ты создан в {CREATION_DATE} командой {', '.join(CREATORS)}.
68
  Ты всегда отвечаешь на русском языке, тепло и профессионально.
@@ -72,6 +84,9 @@ class NeuralChatbot:
72
  Вот вопрос пользователя: """
73
 
74
  def load_model(self):
 
 
 
75
  with st.spinner("🧠 Загружаю нейросеть..."):
76
  try:
77
  from transformers import GPT2LMHeadModel, GPT2Tokenizer
@@ -101,15 +116,12 @@ class NeuralChatbot:
101
  return False
102
 
103
  def generate(self, query):
104
- """Генерация ответа нейросетью"""
105
  if not self.is_loaded:
106
  return self.fallback_response(query)
107
 
108
  try:
109
- # Формируем промпт
110
  prompt = self.system_prompt + query
111
 
112
- # Генерируем
113
  response = self.generator(
114
  prompt,
115
  max_new_tokens=250,
@@ -119,10 +131,8 @@ class NeuralChatbot:
119
  repetition_penalty=1.2
120
  )[0]['generated_text']
121
 
122
- # Убираем промпт
123
  response = response.replace(prompt, "").strip()
124
 
125
- # Если ответ пустой или слишком короткий
126
  if len(response) < 15:
127
  return self.fallback_response(query)
128
 
@@ -133,7 +143,6 @@ class NeuralChatbot:
133
  return self.fallback_response(query)
134
 
135
  def fallback_response(self, query):
136
- """Резервный ответ (если нейросеть не работает)"""
137
  return f"""Я {AI_NAME} от {COMPANY_NAME}.
138
 
139
  К сожалению, сейчас нейросеть временно недоступна, но я хочу ответить на ваш вопрос: "{query}"
@@ -143,7 +152,7 @@ class NeuralChatbot:
143
  А пока я могу рассказать, что создан в {CREATION_DATE} командой {', '.join(CREATORS)}. Я помогаю с научными вопросами и технологиями."""
144
 
145
  # ===================================================================
146
- # 3. ЗАГРУЗКА СТАТЕЙ (ДЛЯ КОНТЕКСТА)
147
  # ===================================================================
148
 
149
  @st.cache_resource
@@ -176,13 +185,16 @@ def load_science_articles():
176
 
177
  @st.cache_resource
178
  def load_embedder():
179
- return SentenceTransformer(EMBEDDING_MODEL)
 
 
 
180
 
181
  @st.cache_resource
182
  def create_embeddings(_articles, _embedder):
183
  if os.path.exists(EMBEDDINGS_FILE):
184
  return np.load(EMBEDDINGS_FILE)
185
- if not _articles:
186
  return np.array([])
187
  texts = [f"{a['title']}\n\n{a['text']}" for a in _articles]
188
  embeddings = _embedder.encode(texts, normalize_embeddings=True, show_progress_bar=True, batch_size=64)
@@ -190,26 +202,27 @@ def create_embeddings(_articles, _embedder):
190
  return embeddings
191
 
192
  def search_articles(query, _articles, _embeddings, _embedder):
193
- """Поиск релевантных статей"""
194
- if not _articles or len(_embeddings) == 0:
 
 
 
 
 
 
 
 
 
 
 
 
195
  return []
196
- query_vector = _embedder.encode([query], normalize_embeddings=True)[0]
197
- scores = _embeddings @ query_vector
198
- top_indices = np.argsort(-scores)[:2]
199
- results = []
200
- for idx in top_indices:
201
- score = float(scores[int(idx)])
202
- if score > 0.15:
203
- article = _articles[int(idx)]
204
- results.append({"title": article['title'], "score": score, "text": article['text'][:500]})
205
- return results
206
 
207
  # ===================================================================
208
- # 4. ОЧИСТКА ЗАПРОСОВ
209
  # ===================================================================
210
 
211
  def clean_query(query):
212
- """Очищает запрос от спама"""
213
  query = re.sub(r'http[s]?://\S+', '', query)
214
  query = re.sub(r'\S+@\S+', '', query)
215
  query = re.sub(r'\+7\s*\(?\d{3}\)?\s*\d{3}\s*\d{2}\s*\d{2}', '', query)
@@ -222,20 +235,8 @@ def clean_query(query):
222
 
223
  return query.strip()
224
 
225
- def enhance_with_context(query, articles_context):
226
- """Добавляет контекст из статей к запросу"""
227
- if not articles_context:
228
- return query
229
-
230
- context = "\n\nВот релевантная научная информация:\n"
231
- for i, art in enumerate(articles_context, 1):
232
- context += f"{i}. {art['title']}\n{art['text'][:300]}...\n"
233
-
234
- context += f"\nНа основе этой информации, ответь на вопрос: {query}"
235
- return context
236
-
237
  # ===================================================================
238
- # 5. ОСНОВНОЙ КЛАСС
239
  # ===================================================================
240
 
241
  class OpenAirAI:
@@ -250,21 +251,12 @@ class OpenAirAI:
250
  self.is_ready = self.chatbot.load_model()
251
  return self.is_ready
252
 
253
- def generate_answer(self, query, articles_context=None):
254
- """Генерирует ответ ТОЛЬКО нейросетью, без if/else"""
255
  clean_q = clean_query(query)
256
-
257
- # Если есть контекст статей - добавляем его
258
- if articles_context:
259
- enhanced_query = enhance_with_context(clean_q, articles_context)
260
- else:
261
- enhanced_query = clean_q
262
-
263
- # ВСЕГДА генерируем нейросетью
264
- return self.chatbot.generate(enhanced_query)
265
 
266
  # ===================================================================
267
- # 6. ИНТЕРФЕЙС
268
  # ===================================================================
269
 
270
  # Загрузка данных
@@ -282,8 +274,7 @@ ai = st.session_state.ai
282
  # История чата
283
  if "messages" not in st.session_state:
284
  st.session_state.messages = []
285
- # Первое приветствие генерируется нейросетью
286
- greeting = ai.generate_answer("Привет! Представься и расскажи о себе")
287
  st.session_state.messages.append({"role": "assistant", "content": greeting})
288
 
289
  # --- БОКОВАЯ ПАНЕЛЬ ---
@@ -311,7 +302,7 @@ with st.sidebar:
311
 
312
  if st.button("🗑️ Очистить чат"):
313
  st.session_state.messages = []
314
- greeting = ai.generate_answer("Привет! Представься и расскажи о себе")
315
  st.session_state.messages.append({"role": "assistant", "content": greeting})
316
  st.rerun()
317
 
@@ -332,21 +323,19 @@ for message in st.session_state.messages:
332
 
333
  # Поле ввода
334
  if prompt := st.chat_input("Задайте вопрос..."):
335
- # Добавляем сообщение пользователя
336
  st.session_state.messages.append({"role": "user", "content": prompt})
337
  with st.chat_message("user"):
338
  st.markdown(prompt)
339
 
340
- # Генерация ответа нейросетью
341
  with st.chat_message("assistant"):
342
  with st.spinner("🧠 Нейросеть генерирует ответ..."):
343
- # Ищем релевантные статьи для контекста
344
  articles_context = search_articles(prompt, articles, embeddings, embedder)
345
 
346
- # Генерируем ответ (ВСЕГДА нейросетью)
347
- response = ai.generate_answer(prompt, articles_context)
348
 
349
- # Добавляем статьи в ответ, если они есть и не встроены
350
  if articles_context and len(response) < 50:
351
  response += "\n\n📄 Я нашел релевантные научные статьи:\n"
352
  for i, art in enumerate(articles_context, 1):
 
24
  if torch.cuda.is_available():
25
  print(f"GPU: {torch.cuda.get_device_name(0)}")
26
 
27
+ # Используем модель, которая точно работает
28
+ MODEL_NAME = "sberbank-ai/rugpt3small_based_on_gpt2" # Рабочая модель
29
  EMBEDDING_MODEL = "all-MiniLM-L6-v2"
30
  SCIENCE_DATASET = "RafaelUI/ru_science"
31
  ARTICLE_LIMIT = 50
 
52
  )
53
 
54
  # ===================================================================
55
+ # 2. ПРОВЕРКА УСТАНОВКИ TRANSFORMERS
56
+ # ===================================================================
57
+
58
+ try:
59
+ from transformers import GPT2LMHeadModel, GPT2Tokenizer
60
+ TRANSFORMERS_OK = True
61
+ except ImportError:
62
+ TRANSFORMERS_OK = False
63
+ st.error("❌ Ошибка: transformers не установлен или устарел")
64
+ st.info("Установите: pip install transformers --upgrade")
65
+
66
+ # ===================================================================
67
+ # 3. НАСТОЯЩАЯ НЕЙРОСЕТЬ
68
  # ===================================================================
69
 
70
  class NeuralChatbot:
 
75
  self.generator = None
76
  self.is_loaded = False
77
 
 
78
  self.system_prompt = f"""Ты - {AI_NAME}, дружелюбный научный AI-ассистент от компании {COMPANY_NAME}.
79
  Ты создан в {CREATION_DATE} командой {', '.join(CREATORS)}.
80
  Ты всегда отвечаешь на русском языке, тепло и профессионально.
 
84
  Вот вопрос пользователя: """
85
 
86
  def load_model(self):
87
+ if not TRANSFORMERS_OK:
88
+ return False
89
+
90
  with st.spinner("🧠 Загружаю нейросеть..."):
91
  try:
92
  from transformers import GPT2LMHeadModel, GPT2Tokenizer
 
116
  return False
117
 
118
  def generate(self, query):
 
119
  if not self.is_loaded:
120
  return self.fallback_response(query)
121
 
122
  try:
 
123
  prompt = self.system_prompt + query
124
 
 
125
  response = self.generator(
126
  prompt,
127
  max_new_tokens=250,
 
131
  repetition_penalty=1.2
132
  )[0]['generated_text']
133
 
 
134
  response = response.replace(prompt, "").strip()
135
 
 
136
  if len(response) < 15:
137
  return self.fallback_response(query)
138
 
 
143
  return self.fallback_response(query)
144
 
145
  def fallback_response(self, query):
 
146
  return f"""Я {AI_NAME} от {COMPANY_NAME}.
147
 
148
  К сожалению, сейчас нейросеть временно недоступна, но я хочу ответить на ваш вопрос: "{query}"
 
152
  А пока я могу рассказать, что создан в {CREATION_DATE} командой {', '.join(CREATORS)}. Я помогаю с научными вопросами и технологиями."""
153
 
154
  # ===================================================================
155
+ # 4. ЗАГРУЗКА СТАТЕЙ
156
  # ===================================================================
157
 
158
  @st.cache_resource
 
185
 
186
  @st.cache_resource
187
  def load_embedder():
188
+ try:
189
+ return SentenceTransformer(EMBEDDING_MODEL)
190
+ except:
191
+ return None
192
 
193
  @st.cache_resource
194
  def create_embeddings(_articles, _embedder):
195
  if os.path.exists(EMBEDDINGS_FILE):
196
  return np.load(EMBEDDINGS_FILE)
197
+ if not _articles or _embedder is None:
198
  return np.array([])
199
  texts = [f"{a['title']}\n\n{a['text']}" for a in _articles]
200
  embeddings = _embedder.encode(texts, normalize_embeddings=True, show_progress_bar=True, batch_size=64)
 
202
  return embeddings
203
 
204
  def search_articles(query, _articles, _embeddings, _embedder):
205
+ if not _articles or len(_embeddings) == 0 or _embedder is None:
206
+ return []
207
+ try:
208
+ query_vector = _embedder.encode([query], normalize_embeddings=True)[0]
209
+ scores = _embeddings @ query_vector
210
+ top_indices = np.argsort(-scores)[:2]
211
+ results = []
212
+ for idx in top_indices:
213
+ score = float(scores[int(idx)])
214
+ if score > 0.15:
215
+ article = _articles[int(idx)]
216
+ results.append({"title": article['title'], "score": score, "text": article['text'][:500]})
217
+ return results
218
+ except:
219
  return []
 
 
 
 
 
 
 
 
 
 
220
 
221
  # ===================================================================
222
+ # 5. ОЧИСТКА ЗАПРОСОВ
223
  # ===================================================================
224
 
225
  def clean_query(query):
 
226
  query = re.sub(r'http[s]?://\S+', '', query)
227
  query = re.sub(r'\S+@\S+', '', query)
228
  query = re.sub(r'\+7\s*\(?\d{3}\)?\s*\d{3}\s*\d{2}\s*\d{2}', '', query)
 
235
 
236
  return query.strip()
237
 
 
 
 
 
 
 
 
 
 
 
 
 
238
  # ===================================================================
239
+ # 6. ОСНОВНОЙ КЛАСС
240
  # ===================================================================
241
 
242
  class OpenAirAI:
 
251
  self.is_ready = self.chatbot.load_model()
252
  return self.is_ready
253
 
254
+ def generate_answer(self, query):
 
255
  clean_q = clean_query(query)
256
+ return self.chatbot.generate(clean_q)
 
 
 
 
 
 
 
 
257
 
258
  # ===================================================================
259
+ # 7. ИНТЕРФЕЙС
260
  # ===================================================================
261
 
262
  # Загрузка данных
 
274
  # История чата
275
  if "messages" not in st.session_state:
276
  st.session_state.messages = []
277
+ greeting = ai.generate_answer("Привет! Представься и расскажи о себе кратко")
 
278
  st.session_state.messages.append({"role": "assistant", "content": greeting})
279
 
280
  # --- БОКОВАЯ ПАНЕЛЬ ---
 
302
 
303
  if st.button("🗑️ Очистить чат"):
304
  st.session_state.messages = []
305
+ greeting = ai.generate_answer("Привет! Представься и расскажи о себе кратко")
306
  st.session_state.messages.append({"role": "assistant", "content": greeting})
307
  st.rerun()
308
 
 
323
 
324
  # Поле ввода
325
  if prompt := st.chat_input("Задайте вопрос..."):
 
326
  st.session_state.messages.append({"role": "user", "content": prompt})
327
  with st.chat_message("user"):
328
  st.markdown(prompt)
329
 
 
330
  with st.chat_message("assistant"):
331
  with st.spinner("🧠 Нейросеть генерирует ответ..."):
332
+ # Ищем релева��тные статьи
333
  articles_context = search_articles(prompt, articles, embeddings, embedder)
334
 
335
+ # Генерируем ответ
336
+ response = ai.generate_answer(prompt)
337
 
338
+ # Добавляем статьи в ответ
339
  if articles_context and len(response) < 50:
340
  response += "\n\n📄 Я нашел релевантные научные статьи:\n"
341
  for i, art in enumerate(articles_context, 1):