/
dava
/
MagDipl
Обзор
Документация
Войти
/
dava
/
MagDipl
Код
Запросы
0
Задачи
Вики
Пакеты
0
Релизы
0
CI/CD
Аналитика
Безопасность
master
core/model.py
208 строк
9 KB
david
beta_version
23 фев 2026, 11:20
23 фев 2026, 11:20
777dfb5
Код
Авторство
О чём код?
import pandas as pd import numpy as np from sklearn.feature_extraction.text import TfidfVectorizer from sklearn.naive_bayes import MultinomialNB from sklearn.linear_model import SGDClassifier from sklearn.ensemble import RandomForestClassifier from sklearn.model_selection import train_test_split from sklearn.metrics import accuracy_score, classification_report from core.preprocess import preprocess_text import logging logger = logging.getLogger(__name__) class NewsClassifier: def __init__(self, model_type="svm"): """ Инициализация классификатора. Args: model_type: Тип модели - "nb" (Naive Bayes), "svm" (SVM), "rf" (Random Forest) """ self.vectorizer = TfidfVectorizer( max_features=10000, ngram_range=(1, 2), # Униграммы и биграммы min_df=2, # Минимальная частота слова max_df=0.95 # Максимальная частота слова ) if model_type == "svm": self.model = SGDClassifier( loss='hinge', penalty='l2', alpha=1e-4, random_state=42, max_iter=1000, tol=1e-3 ) elif model_type == "rf": self.model = RandomForestClassifier( n_estimators=100, random_state=42, n_jobs=-1 ) else: self.model = MultinomialNB(alpha=0.1) self.model_type = model_type self.is_trained = False self.categories = [] def train(self, dataset_path="data/dataset.csv", min_samples_per_category=5): """ Обучает классификатор на датасете. Args: dataset_path: Путь к CSV файлу с колонками title, text, category min_samples_per_category: Минимальное количество примеров на категорию Returns: accuracy: Точность модели на тестовой выборке """ logger.info(f"Загрузка датасета из {dataset_path}") df = pd.read_csv(dataset_path) # Проверяем наличие нужных колонок required_cols = ["text", "category"] missing = [col for col in required_cols if col not in df.columns] if missing: raise ValueError(f"Отсутствуют колонки: {missing}") # Удаляем строки с пустым текстом или категорией df = df[df["text"].notna() & (df["text"].str.strip() != "")] df = df[df["category"].notna() & (df["category"].str.strip() != "")] # Проверяем, есть ли размеченные данные (не "unknown") unknown_count = len(df[df["category"] == "unknown"]) labeled_count = len(df[df["category"] != "unknown"]) logger.info(f"Статей с категорией 'unknown': {unknown_count}") logger.info(f"Размеченных статей: {labeled_count}") if labeled_count == 0: raise ValueError( "В датасете нет размеченных данных для обучения!\n\n" "Все статьи имеют категорию 'unknown'. Для обучения классификатора нужны размеченные данные.\n\n" "Варианты решения:\n" "1. Создайте обучающий датасет: вручную пометьте несколько статей категориями (politics, sport, economics и т.д.)\n" "2. Используйте предобученную модель (если доступна)\n" "3. Используйте другой датасет с размеченными данными" ) # Удаляем категорию "unknown" если она есть (для обучения нужны размеченные данные) df = df[df["category"] != "unknown"] if len(df) == 0: raise ValueError("Датасет пуст после фильтрации") # Фильтруем категории с малым количеством примеров category_counts = df["category"].value_counts() valid_categories = category_counts[category_counts >= min_samples_per_category].index df = df[df["category"].isin(valid_categories)] if len(df) == 0: raise ValueError(f"Нет категорий с минимум {min_samples_per_category} примерами") logger.info(f"Категории для обучения: {sorted(valid_categories.tolist())}") logger.info(f"Распределение по категориям:\n{category_counts[valid_categories]}") df["processed"] = df["text"].apply(preprocess_text) X = df["processed"] y = df["category"] self.categories = sorted(y.unique().tolist()) X_train, X_test, y_train, y_test = train_test_split( X, y, test_size=0.2, random_state=42, stratify=y ) logger.info(f"Обучающая выборка: {len(X_train)}, тестовая: {len(X_test)}") X_train_vec = self.vectorizer.fit_transform(X_train) X_test_vec = self.vectorizer.transform(X_test) logger.info(f"Обучение модели {self.model_type}...") self.model.fit(X_train_vec, y_train) predictions = self.model.predict(X_test_vec) accuracy = accuracy_score(y_test, predictions) # Выводим детальный отчёт report = classification_report(y_test, predictions, target_names=self.categories) logger.info(f"Classification Report:\n{report}") self.is_trained = True return accuracy def predict(self, text: str): """ Предсказывает категорию для текста. Args: text: Текст статьи Returns: str: Предсказанная категория """ if not self.is_trained: return "Модель не обучена" processed = preprocess_text(text) vectorized = self.vectorizer.transform([processed]) prediction = self.model.predict(vectorized) return prediction[0] def classify_articles(self, articles: list) -> list: """ Классифицирует список статей. Args: articles: Список словарей с ключами "title", "text", "category" (опционально) Returns: list: Список статей с обновлённым полем "category" """ if not self.is_trained: logger.warning("Модель не обучена, возвращаю статьи без классификации") return articles texts = [article.get("text", "") for article in articles] processed_texts = [preprocess_text(text) for text in texts] if not processed_texts: return articles vectorized = self.vectorizer.transform(processed_texts) predictions = self.model.predict(vectorized) # Обновляем категории в статьях for i, article in enumerate(articles): article["category"] = predictions[i] logger.info(f"Классифицировано {len(articles)} статей") return articles def predict_topk(self, text: str, k: int = 3): """ Возвращает топ-k категорий с вероятностями (если модель это поддерживает). """ if not self.is_trained: return [] processed = preprocess_text(text) vectorized = self.vectorizer.transform([processed]) if not hasattr(self.model, "predict_proba"): label = self.model.predict(vectorized)[0] return [(label, 1.0)] probs = self.model.predict_proba(vectorized)[0] classes = list(self.model.classes_) pairs = list(zip(classes, probs)) pairs.sort(key=lambda x: x[1], reverse=True) return pairs[: max(1, int(k))]