/
dakone22
/
ppa
Обзор
Документация
Войти
/
dakone22
/
ppa
Код
Запросы
0
Задачи
Вики
Пакеты
0
Релизы
0
CI/CD
Аналитика
Безопасность
master
lab2/src/prepare_data.py
311 строк
13 KB
dakone22
created: lab2, lab3, lab4
04 июн 2026, 14:19
Верифицирован
04 июн 2026, 14:19
03322e8
Код
Авторство
О чём код?
""" Скрипт для подготовки данных сообщений Telegram с поддержкой чекпоинтов. Запуск: - Полный: python prepare_data.py - Тест (500 сообщений): python prepare_data.py --test - Лимит (N сообщений): python prepare_data.py --limit 1000 Возобновление: при повторном запуске скрипт продолжит с места остановки. """ import argparse import json from pathlib import Path import pandas as pd import torch import numpy as np from tqdm import tqdm from transformers import AutoModelForTokenClassification, AutoTokenizer from sentence_transformers import SentenceTransformer # Конфигурация BASE_DIR = Path(__file__).parent DATA_DIR = BASE_DIR / "data" OUTPUT_PREPARED = BASE_DIR / "prepared_messages.parquet" OUTPUT_FEATURES = BASE_DIR / "messages_with_features.parquet" CHECKPOINT_DIR = BASE_DIR / "checkpoints" CHECKPOINT_NER = CHECKPOINT_DIR / "ner_checkpoint.parquet" CHECKPOINT_EMB = CHECKPOINT_DIR / "emb_checkpoint.parquet" # Модели NER_MODEL = "Babelscape/wikineural-multilingual-ner" EMB_MODEL = "cointegrated/rubert-tiny2" # Параметры MIN_TEXT_LEN = 15 BATCH_SIZE_NER = 64 # Увеличено для CUDA BATCH_SIZE_EMB = 256 # Увеличено для CUDA CHECKPOINT_INTERVAL = 10 # Сохранять чекпоинт каждые N сообщений def load_all_messages(limit=None): """Загружает сообщения из всех каналов.""" print("Загрузка сообщений из JSONL файлов...") all_data = [] channels = [d.name for d in DATA_DIR.iterdir() if d.is_dir()] print(f"Найдено каналов: {channels}") for channel in tqdm(channels, desc="Каналы"): jsonl_path = DATA_DIR / channel / "all_messages.jsonl" if not jsonl_path.exists(): continue with open(jsonl_path, 'r', encoding='utf-8') as f: for line in f: try: msg = json.loads(line.strip()) record = { 'msg_id': msg.get('message_id'), 'channel': msg.get('channel', channel), 'text': msg.get('text', ''), 'timestamp': msg.get('date'), 'views': msg.get('views', 0), 'forwards': msg.get('forwards', 0), 'replies': msg.get('replies', 0) } all_data.append(record) if limit and len(all_data) >= limit: break except json.JSONDecodeError: continue if limit and len(all_data) >= limit: break print(f"Загружено записей: {len(all_data)}") return all_data def clean_data(records): """Очистка данных.""" print("Очистка данных...") df = pd.DataFrame(records) df = df.drop_duplicates(subset=['msg_id', 'channel']) df['text_len'] = df['text'].apply(lambda x: len(str(x)) if x else 0) initial_count = len(df) df = df[df['text_len'] >= MIN_TEXT_LEN].copy() print(f"Удалено {initial_count - len(df)} записей (длина < {MIN_TEXT_LEN})") df['timestamp'] = pd.to_datetime(df['timestamp']) return df def init_ner_model(): """Инициализация NER модели для прямого использования.""" print("Инициализация NER модели...") device = torch.device("cuda" if torch.cuda.is_available() else "cpu") print(f"Используется устройство: {device}") tokenizer = AutoTokenizer.from_pretrained(NER_MODEL) model = AutoModelForTokenClassification.from_pretrained(NER_MODEL) model.to(device) model.eval() return model, tokenizer, device def extract_entities_with_model(texts, model, tokenizer, device, start_idx=0, checkpoint_data=None): """Извлечение сущностей с использованием модели напрямую.""" entities_list = checkpoint_data if checkpoint_data else [] num_texts = len(texts) # Если есть чекпоинт, начинаем с соответствующего индекса if start_idx > 0: print(f"Возобновление NER с индекса {start_idx}") with torch.no_grad(): for i in tqdm(range(start_idx, num_texts), desc="NER", initial=start_idx, total=num_texts): text = texts[i] try: # Токенизация inputs = tokenizer(text, return_tensors="pt", truncation=True, max_length=512, padding=True) inputs = {k: v.to(device) for k, v in inputs.items()} # Предсказание outputs = model(**inputs) predictions = torch.argmax(outputs.logits, dim=2) # Преобразование в сущности tokens = tokenizer.convert_ids_to_tokens(inputs["input_ids"][0]) entities = [] current_entity = None for token, pred_id in zip(tokens, predictions[0]): if token in ['[CLS]', '[SEP]', '[PAD]']: continue # Получаем метку pred_label = model.config.id2label[pred_id.item()] if pred_label.startswith('B-'): if current_entity: entities.append(current_entity) entity_type = pred_label[2:] current_entity = {'word': token.replace('##', ''), 'type': entity_type} elif pred_label.startswith('I-') and current_entity: current_entity['word'] += token.replace('##', '') else: if current_entity: entities.append(current_entity) current_entity = None if current_entity: entities.append(current_entity) # Фильтрация по типам (PER, ORG, LOC) filtered = [ent['word'] for ent in entities if ent['type'] in ['PER', 'ORG', 'LOC']] entities_list.append(filtered) except Exception as e: print(f"Ошибка NER для текста {i}: {e}") entities_list.append([]) # Сохранение чекпоинта if (i + 1) % CHECKPOINT_INTERVAL == 0: checkpoint_df = pd.DataFrame({'entities': entities_list}) checkpoint_df.to_parquet(CHECKPOINT_NER, index=False) return entities_list def compute_embeddings_with_checkpoint(texts, model, start_idx=0, checkpoint_data=None): """Вычисление эмбеддингов с чекпоинтами.""" embeddings = checkpoint_data if checkpoint_data else [] num_texts = len(texts) if start_idx > 0: print(f"Возобновление эмбеддингов с индекса {start_idx}") num_batches = (num_texts + BATCH_SIZE_EMB - 1) // BATCH_SIZE_EMB start_batch = start_idx // BATCH_SIZE_EMB for i in tqdm(range(start_idx, num_texts, BATCH_SIZE_EMB), desc="Embeddings", initial=start_batch, total=num_batches): batch = texts[i:i + BATCH_SIZE_EMB] try: with torch.no_grad(): emb = model.encode(batch, convert_to_numpy=True, show_progress_bar=False) embeddings.extend(emb.tolist()) except Exception as e: print(f"Ошибка эмбеддинга для батча с индексом {i}: {e}") dim = model.get_sentence_embedding_dimension() embeddings.extend([[0.0] * dim for _ in range(len(batch))]) # Сохранение чекпоинта if (i + BATCH_SIZE_EMB) % CHECKPOINT_INTERVAL == 0 or i + BATCH_SIZE_EMB >= num_texts: checkpoint_df = pd.DataFrame({'embedding': embeddings}) checkpoint_df.to_parquet(CHECKPOINT_EMB, index=False) return embeddings def load_checkpoint(checkpoint_path): """Загрузка чекпоинта.""" if checkpoint_path.exists(): print(f"Найден чекпоинт: {checkpoint_path}") df = pd.read_parquet(checkpoint_path) return df return None def main(): parser = argparse.ArgumentParser(description="Подготовка данных Telegram.") parser.add_argument('--test', action='store_true', help='Тестовый запуск (500 сообщений)') parser.add_argument('--limit', type=int, default=None, help='Ограничение числа сообщений') parser.add_argument('--force-restart', action='store_true', help='Принудительный перезапуск (игнор чекпоинтов)') args = parser.parse_args() # Создаем папку для чекпоинтов CHECKPOINT_DIR.mkdir(exist_ok=True) # Определяем устройство device_emb = 'cuda' if torch.cuda.is_available() else 'cpu' print(f"Устройство для эмбеддингов: {device_emb}") # 1. Загрузка и очистка (или загрузка из сохраненного) if OUTPUT_PREPARED.exists() and not args.force_restart: print(f"Загрузка подготовленных данных из {OUTPUT_PREPARED}") df = pd.read_parquet(OUTPUT_PREPARED) print(f"Загружено {len(df)} записей") else: records = load_all_messages(limit=500 if args.test else args.limit) if not records: print("Нет данных") return df = clean_data(records) df.to_parquet(OUTPUT_PREPARED, index=False) print(f"Сохранен промежуточный файл: {OUTPUT_PREPARED}") texts = df['text'].tolist() # 2. NER с чекпоинтами ner_checkpoint = None start_idx_ner = 0 if CHECKPOINT_NER.exists() and not args.force_restart: ner_checkpoint_df = load_checkpoint(CHECKPOINT_NER) if ner_checkpoint_df is not None and 'entities' in ner_checkpoint_df.columns: ner_checkpoint = ner_checkpoint_df['entities'].tolist() start_idx_ner = len(ner_checkpoint) print(f"Восстановлено {start_idx_ner} entities из чекпоинта") if start_idx_ner < len(texts): model_ner, tokenizer, device_ner = init_ner_model() entities = extract_entities_with_model( texts, model_ner, tokenizer, device_ner, start_idx=start_idx_ner, checkpoint_data=ner_checkpoint ) df['entities'] = entities # Удаляем чекпоинт после успешного завершения if CHECKPOINT_NER.exists(): CHECKPOINT_NER.unlink() print("Чекпоинт NER удален") else: print("NER уже выполнен полностью (из чекпоинта)") if ner_checkpoint: df['entities'] = ner_checkpoint df['shared_entities_count'] = 0 # Заглушка # 3. Эмбеддинги с чекпоинтами emb_checkpoint = None start_idx_emb = 0 if CHECKPOINT_EMB.exists() and not args.force_restart: emb_checkpoint_df = load_checkpoint(CHECKPOINT_EMB) if emb_checkpoint_df is not None and 'embedding' in emb_checkpoint_df.columns: emb_checkpoint = emb_checkpoint_df['embedding'].tolist() start_idx_emb = len(emb_checkpoint) print(f"Восстановлено {start_idx_emb} embeddings из чекпоинта") if start_idx_emb < len(texts): print("\nИнициализация модели эмбеддингов...") emb_model = SentenceTransformer(EMB_MODEL, device=device_emb) embeddings = compute_embeddings_with_checkpoint( texts, emb_model, start_idx=start_idx_emb, checkpoint_data=emb_checkpoint ) df['embedding'] = embeddings # Удаляем чекпоинт после успешного завершения if CHECKPOINT_EMB.exists(): CHECKPOINT_EMB.unlink() print("Чекпоинт эмбеддингов удален") else: print("Эмбеддинги уже вычислены (из чекпоинта)") if emb_checkpoint: df['embedding'] = emb_checkpoint # 4. Сохранение финального результата df.to_parquet(OUTPUT_FEATURES, index=False) print(f"\nФинальный файл: {OUTPUT_FEATURES}") print(f"Обработано: {len(df)} сообщений") print(f"Признаки: entities, shared_entities_count, embedding") if __name__ == "__main__": main()