OpenAirAI / app.py
X
Update app.py
efceacb verified
Raw History Blame
18.2 kB
import streamlit as st
from datasets import load_dataset
import numpy as np
from sentence_transformers import SentenceTransformer
import time
from datetime import datetime
import json
import os
import pandas as pd
import pickle
import random
import torch
import torch.nn as nn
import torch.nn.functional as F
from transformers import AutoTokenizer, AutoModel
from torch.optim import AdamW
import gc
import warnings
warnings.filterwarnings('ignore')
# ===================================================================
# 1. НАСТРОЙКИ (ОПТИМИЗИРОВАННЫЕ)
# ===================================================================
# Проверка CUDA
print(f"CUDA доступна: {torch.cuda.is_available()}")
if torch.cuda.is_available():
print(f"Количество GPU: {torch.cuda.device_count()}")
print(f"GPU: {torch.cuda.get_device_name(0)}")
os.environ["CUDA_VISIBLE_DEVICES"] = "0"
MODEL_NAME = "DeepPavlov/rubert-base-cased" # Для русского языка
EMBEDDING_MODEL = "all-MiniLM-L6-v2" # ЛЕГКАЯ И БЫСТРАЯ модель для эмбеддингов
# Если нужно точнее, можно использовать: "BAAI/bge-large-en-v1.5" (медленнее)
SCIENCE_DATASET = "RafaelUI/ru_science"
ARTICLE_LIMIT = 100 # БЫСТРЫЙ СТАРТ - 100 статей (можно увеличить позже)
MAX_LENGTH = 256 # Уменьшаем для скорости
BATCH_SIZE = 32 # Увеличиваем для скорости
EPOCHS = 2
ADMIN_USER = "admin"
ADMIN_PASS = "hfpassword21"
LOG_FILE = "query_logs.json"
EMBEDDINGS_FILE = "science_embeddings.npy"
ARTICLES_FILE = "science_articles.pkl"
DIALOG_MODEL_PATH = "openairai_dialog_model.bin"
# Информация о создателях
AI_NAME = "OpenAirAI"
COMPANY_NAME = "OpenRussianAI"
CREATORS = ["Грибков Евгений", "RootLinux21"]
WEBSITE = "https://sites.google.com/view/opruai/home"
HUGGINGFACE = "https://huggingface.co/OpenRussianAI"
CREATION_DATE = "2026"
# Обучающие диалоги
TRAINING_DIALOGS = [
{
"context": "Привет",
"response": "Привет! Я OpenAirAI — ваш научный AI-ассистент от OpenRussianAI. Создан в 2026 году для работы с научными статьями. Чем могу помочь?"
},
{
"context": "Кто ты",
"response": "Меня зовут OpenAirAI. Я — AI-ассистент, созданный компанией OpenRussianAI в 2026 году командой разработчиков Грибков Евгений и RootLinux21."
},
{
"context": "Кто тебя создал",
"response": "Меня создала команда OpenRussianAI в составе Грибкова Евгения и RootLinux21 в 2026 году."
},
{
"context": "Что ты умеешь",
"response": "Я умею анализировать научные статьи, находить информацию, помогать с исследованиями в области сельского хозяйства, биологии, химии."
},
{
"context": "Где ваш сайт",
"response": f"Сайт OpenRussianAI: {WEBSITE}"
},
{
"context": "Где ваши модели",
"response": f"Модели OpenRussianAI на Hugging Face: {HUGGINGFACE}"
},
{
"context": "Спасибо",
"response": "Всегда рад помочь! Я, OpenAirAI, здесь для вас. Обращайтесь! 😊"
},
{
"context": "Пока",
"response": "До свидания! OpenAirAI всегда на связи. Удачи в исследованиях! 👋"
}
]
# Настройка страницы
st.set_page_config(
page_title=f"{AI_NAME} - Научный AI-ассистент",
page_icon="🧪",
layout="wide",
initial_sidebar_state="expanded"
)
# Загружаем логи
if os.path.exists(LOG_FILE):
with open(LOG_FILE, "r") as f:
query_logs = json.load(f)
else:
query_logs = []
# ===================================================================
# 2. МОДЕЛЬ ДЛЯ ДИАЛОГОВ
# ===================================================================
class DialogModel(nn.Module):
def __init__(self, pretrained_name, num_labels=2):
super().__init__()
self.bert = AutoModel.from_pretrained(pretrained_name)
self.classifier = nn.Linear(self.bert.config.hidden_size, num_labels)
self.dropout = nn.Dropout(0.1)
def forward(self, input_ids, attention_mask):
outputs = self.bert(input_ids=input_ids, attention_mask=attention_mask)
pooled = outputs.pooler_output
pooled = self.dropout(pooled)
logits = self.classifier(pooled)
return logits
class OpenAirAI:
def __init__(self):
self.name = AI_NAME
self.company = COMPANY_NAME
self.creators = CREATORS
self.creation_date = CREATION_DATE
self.device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
self.tokenizer = AutoTokenizer.from_pretrained(MODEL_NAME)
self.model = None
self.is_trained = False
self.model_path = DIALOG_MODEL_PATH
self.contexts = [d["context"] for d in TRAINING_DIALOGS]
self.responses = [d["response"] for d in TRAINING_DIALOGS]
def train_on_dialogs(self):
with st.spinner(f"🧠 Обучаю {self.name} на диалогах..."):
self.model = DialogModel(MODEL_NAME).to(self.device)
optimizer = AdamW(self.model.parameters(), lr=2e-5)
self.model.train()
progress_bar = st.progress(0)
for epoch in range(EPOCHS):
total_loss = 0
for i in range(0, len(self.contexts), BATCH_SIZE):
batch_contexts = self.contexts[i:i+BATCH_SIZE]
encodings = self.tokenizer(
batch_contexts,
truncation=True,
padding=True,
max_length=MAX_LENGTH,
return_tensors='pt'
)
input_ids = encodings['input_ids'].to(self.device)
attention_mask = encodings['attention_mask'].to(self.device)
labels = torch.tensor([1] * len(batch_contexts)).to(self.device)
optimizer.zero_grad()
logits = self.model(input_ids, attention_mask)
loss = F.cross_entropy(logits, labels)
loss.backward()
optimizer.step()
total_loss += loss.item()
progress_bar.progress((epoch + i/len(self.contexts)) / EPOCHS)
st.write(f"Эпоха {epoch+1}/{EPOCHS}, Потери: {total_loss/len(self.contexts):.4f}")
torch.save(self.model.state_dict(), self.model_path)
self.is_trained = True
st.success(f"✅ {self.name} обучен!")
def load_model(self):
if os.path.exists(self.model_path):
try:
self.model = DialogModel(MODEL_NAME).to(self.device)
self.model.load_state_dict(torch.load(self.model_path, map_location=self.device))
self.model.eval()
self.is_trained = True
return True
except:
return False
return False
def generate_response(self, query):
if not self.is_trained:
return f"Я {self.name}, научный AI-ассистент от {self.company}. Чем могу помочь?"
self.model.eval()
with torch.no_grad():
encodings = self.tokenizer(
query,
truncation=True,
padding=True,
max_length=MAX_LENGTH,
return_tensors='pt'
)
input_ids = encodings['input_ids'].to(self.device)
attention_mask = encodings['attention_mask'].to(self.device)
outputs = self.model.bert(input_ids=input_ids, attention_mask=attention_mask)
query_embedding = outputs.pooler_output
context_embeddings = []
for context in self.contexts:
ctx_enc = self.tokenizer(
context,
truncation=True,
padding=True,
max_length=MAX_LENGTH,
return_tensors='pt'
)
ctx_input_ids = ctx_enc['input_ids'].to(self.device)
ctx_attention_mask = ctx_enc['attention_mask'].to(self.device)
ctx_outputs = self.model.bert(input_ids=ctx_input_ids, attention_mask=ctx_attention_mask)
context_embeddings.append(ctx_outputs.pooler_output)
context_embeddings = torch.cat(context_embeddings, dim=0)
similarities = F.cosine_similarity(query_embedding, context_embeddings)
best_idx = torch.argmax(similarities).item()
if similarities[best_idx] > 0.5:
return self.responses[best_idx]
else:
return f"Я {self.name}, научный AI-ассистент от {self.company}. Создан в {self.creation_date}. Чем могу помочь?"
# ===================================================================
# 3. БЫСТРАЯ ЗАГРУЗКА НАУЧНЫХ СТАТЕЙ
# ===================================================================
@st.cache_resource
def load_science_articles():
articles_file = ARTICLES_FILE
if os.path.exists(articles_file):
with st.spinner("📚 Загружаю научные статьи с диска..."):
with open(articles_file, 'rb') as f:
return pickle.load(f)
with st.spinner("📚 Загружаю научные статьи (первый раз, ~1-2 минуты)..."):
try:
dataset = load_dataset(SCIENCE_DATASET, split="train", streaming=True)
articles = []
for i, row in enumerate(dataset):
if i >= ARTICLE_LIMIT:
break
text = row.get('content', '') or row.get('text', '') or str(row)
title = row.get('title', f"Статья {i}")
articles.append({
"id": i,
"title": title[:200],
"text": text[:2000], # Уменьшаем для скорости
"source": "ru_science"
})
with open(articles_file, 'wb') as f:
pickle.dump(articles, f)
return articles
except Exception as e:
st.error(f"Ошибка: {e}")
return create_test_articles()
def create_test_articles():
return [
{
"id": 1,
"title": "Влияние удобрений на рост растений",
"text": "Исследование показывает, что применение азотных удобрений увеличивает урожайность.",
"source": "test"
},
{
"id": 2,
"title": "Методы биоконверсии",
"text": "Биоконверсия позволяет перерабатывать органические отходы в удобрения.",
"source": "test"
}
]
@st.cache_resource
def load_embedder():
with st.spinner("🧠 Загружаю модель для эмбеддингов..."):
return SentenceTransformer(EMBEDDING_MODEL)
@st.cache_resource
def create_embeddings(_articles, _embedder):
embeddings_file = EMBEDDINGS_FILE
if os.path.exists(embeddings_file):
with st.spinner("📊 Загружаю эмбеддинги с диска..."):
return np.load(embeddings_file)
with st.spinner(f"🔢 Создаю эмбеддинги для {len(_articles)} статей (1-2 минуты)..."):
texts = [f"{a['title']}\n\n{a['text']}" for a in _articles]
embeddings = _embedder.encode(
texts,
normalize_embeddings=True,
show_progress_bar=True,
batch_size=64, # Увеличен для скорости
device='cuda' if torch.cuda.is_available() else 'cpu'
)
np.save(embeddings_file, embeddings)
return embeddings
# ===================================================================
# 4. ФУНКЦИИ ПОИСКА
# ===================================================================
def search_science(query, _articles, _embeddings, _embedder):
if not query:
return None
start_time = time.time()
query_vector = _embedder.encode([query], normalize_embeddings=True)[0]
scores = _embeddings @ query_vector
top_indices = np.argsort(-scores)[:3]
results = []
for idx in top_indices:
score = float(scores[int(idx)])
if score > 0.2:
article = _articles[int(idx)]
results.append({
"title": article['title'],
"score": score,
"text": article['text'][:1000],
"source": article.get('source', 'ru_science')
})
log_entry = {
"timestamp": datetime.now().isoformat(),
"query": query,
"results_count": len(results),
"response_time": round(time.time() - start_time, 2)
}
query_logs.append(log_entry)
with open(LOG_FILE, "w") as f:
json.dump(query_logs[-100:], f)
return results
def clear_cache():
files = [EMBEDDINGS_FILE, ARTICLES_FILE, DIALOG_MODEL_PATH]
for file in files:
if os.path.exists(file):
os.remove(file)
st.cache_resource.clear()
return True
# ===================================================================
# 5. ИНТЕРФЕЙС
# ===================================================================
# ЗАГРУЗКА (быстрая)
articles = load_science_articles()
embedder = load_embedder()
embeddings = create_embeddings(articles, embedder)
# Инициализация AI
if 'dialog_ai' not in st.session_state:
st.session_state.dialog_ai = OpenAirAI()
if not st.session_state.dialog_ai.load_model():
st.session_state.dialog_ai.train_on_dialogs()
dialog_ai = st.session_state.dialog_ai
# --- БОКОВАЯ ПАНЕЛЬ ---
with st.sidebar:
st.image("https://cdn-icons-png.flaticon.com/512/4248/4248455.png", width=80)
st.title(f"👑 {AI_NAME}")
st.markdown(f"""
**{dialog_ai.name}** | {dialog_ai.creation_date}
**Компания:** {dialog_ai.company}
**Разработчики:** {', '.join(dialog_ai.creators)}
[🌐 Сайт]({WEBSITE})
[🤗 HF]({HUGGINGFACE})
""")
st.divider()
# Админка
if "logged_in" not in st.session_state:
st.session_state.logged_in = False
if not st.session_state.logged_in:
with st.form("login_form"):
username = st.text_input("👤 Логин", placeholder="admin")
password = st.text_input("🔑 Пароль", type="password", placeholder="hfpassword21")
if st.form_submit_button("🔑 Войти"):
if username == ADMIN_USER and password == ADMIN_PASS:
st.session_state.logged_in = True
st.rerun()
else:
st.error("❌ Неверно")
else:
st.success("✅ Админ")
if st.button("🚪 Выйти"):
st.session_state.logged_in = False
st.rerun()
if st.button("🔄 Переобучить AI"):
dialog_ai.train_on_dialogs()
st.rerun()
if st.button("🗑️ Очистить кэш"):
clear_cache()
st.success("Кэш очищен!")
st.rerun()
# Статистика
st.metric("Всего запросов", len(query_logs))
st.metric("Статей", len(articles))
if os.path.exists(EMBEDDINGS_FILE):
size = os.path.getsize(EMBEDDINGS_FILE) / (1024 * 1024)
st.metric("Эмбеддинги", f"{size:.1f} MB")
# --- ОСНОВНАЯ ЧАСТЬ ---
st.title(f"🧪 {AI_NAME} - Научный AI-ассистент")
st.markdown(f"**{AI_NAME}** от **{COMPANY_NAME}** | Работает с научными статьями")
# Приветствие
if "greeting_shown" not in st.session_state:
st.session_state.greeting_shown = True
st.success(f"🤖 **{AI_NAME}:** {dialog_ai.generate_response('Привет')}")
st.info(f"📚 {len(articles)} научных статей загружено")
# Поиск
query = st.text_input(
"🔍 Что хочешь узнать?",
placeholder="Например: Как удобрения влияют на урожайность?",
key="query_input"
)
if query:
with st.spinner("🔎 Ищу..."):
results = search_science(query, articles, embeddings, embedder)
if results:
for i, result in enumerate(results, 1):
with st.expander(f"#{i} {result['title']} (сходство: {result['score']:.2f})", expanded=i==1):
st.write(result['text'] + "...")
st.caption(f"📌 {result['source']}")
else:
st.warning("😕 Не нашёл подходящих статей")
else:
st.info("💡 Напиши вопрос о науке")
# Подвал
st.divider()
st.caption(f"🧪 {AI_NAME} от {COMPANY_NAME} | {CREATION_DATE} | ru_science")