diff --git a/jupyter_notebooks/.ipynb_checkpoints/Section_2-3-4_Figure_8_Step_1_LLM_annotation-checkpoint.ipynb b/jupyter_notebooks/.ipynb_checkpoints/Section_2-3-4_Figure_8_Step_1_LLM_annotation-checkpoint.ipynb new file mode 100644 index 0000000000000000000000000000000000000000..f175a91888eac9bb2f3101d96991e9f96042afca --- /dev/null +++ b/jupyter_notebooks/.ipynb_checkpoints/Section_2-3-4_Figure_8_Step_1_LLM_annotation-checkpoint.ipynb @@ -0,0 +1,1451 @@ +{ + "cells": [ + { + "cell_type": "markdown", + "id": "23d0ae58", + "metadata": {}, + "source": [ + "# Deepfake Adapter Dataset - LLM Annotation Pipeline" + ] + }, + { + "cell_type": "markdown", + "id": "e4407358", + "metadata": {}, + "source": [ + "### Unified Model Loading & Inference\n", + "Code for querying Mistral, Gemma, and Qwen models." + ] + }, + { + "cell_type": "markdown", + "id": "1a1b9d0e", + "metadata": {}, + "source": [ + "## CLEANING & PREPROCESSING" + ] + }, + { + "cell_type": "markdown", + "id": "3df42c46", + "metadata": {}, + "source": [ + "#### Named Entity Recognitition (NER) using SpaCy " + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "a287eef4", + "metadata": {}, + "outputs": [], + "source": [ + "import pandas as pd\n", + "import re\n", + "from pathlib import Path\n", + "import emoji\n", + "import spacy\n", + "\n", + "# Load spaCy model\n", + "# You may need to download it first: python -m spacy download en_core_web_sm\n", + "try:\n", + " nlp = spacy.load(\"en_core_web_sm\")\n", + " print(\"✅ spaCy model loaded: en_core_web_sm\")\n", + "except OSError:\n", + " print(\"❌ spaCy model not found. Downloading...\")\n", + " import subprocess\n", + " subprocess.run([\"python\", \"-m\", \"spacy\", \"download\", \"en_core_web_sm\"])\n", + " nlp = spacy.load(\"en_core_web_sm\")\n", + " print(\"✅ spaCy model downloaded and loaded\")\n", + "\n", + "# Set up paths\n", + "current_dir = Path.cwd()\n", + "#input_file = current_dir.parent / \"data/CSV/real_person_adapters.csv\"\n", + "input_file = current_dir.parent / \"data/CSV/model_adapter/real_person_adapter.csv\"\n", + "\n", + "# Load dataset\n", + "df = pd.read_csv(input_file)\n", + "print(f\"Loaded {len(df)} rows\")\n", + "\n", + "def translate_leetspeak(text: str) -> str:\n", + " \"\"\"\n", + " Translate common leetspeak patterns to normal letters.\n", + " Examples: 4kira -> Akira, 3mma -> Emma, 1rene -> Irene\n", + " \"\"\"\n", + " if not text:\n", + " return text\n", + " \n", + " # Common leetspeak mappings (order matters!)\n", + " leetspeak_map = {\n", + " '4': 'a',\n", + " '3': 'e', \n", + " '1': 'i',\n", + " '0': 'o',\n", + " '7': 't',\n", + " '5': 's',\n", + " '8': 'b',\n", + " '9': 'g',\n", + " '@': 'a',\n", + " '$': 's',\n", + " '!': 'i',\n", + " }\n", + " \n", + " result = text\n", + " # Apply mappings at word boundaries or start of string\n", + " for leet, normal in leetspeak_map.items():\n", + " # Replace at start of word\n", + " result = re.sub(rf'\\b{re.escape(leet)}', normal, result, flags=re.IGNORECASE)\n", + " # Replace standalone numbers that look like letters in context\n", + " result = re.sub(rf'(?<=[a-z]){re.escape(leet)}(?=[a-z])', normal, result, flags=re.IGNORECASE)\n", + " \n", + " return result\n", + "\n", + "def preprocess_for_ner(name: str) -> str:\n", + " \"\"\"\n", + " Preprocess the name before spaCy NER.\n", + " Remove noise but keep the actual name parts.\n", + " \"\"\"\n", + " if pd.isna(name):\n", + " return \"\"\n", + " \n", + " name = str(name)\n", + " \n", + " # FIRST: Translate leetspeak\n", + " name = translate_leetspeak(name)\n", + " \n", + " # Remove emoji\n", + " name = emoji.replace_emoji(name, replace=' ')\n", + " \n", + " # Remove version indicators (v1, v2, v1.0, etc.)\n", + " name = re.sub(r'\\s*[vV]\\d+(\\.\\d+)?\\s*', ' ', name)\n", + " \n", + " # Remove LoRA-related terms (case insensitive)\n", + " lora_terms = ['lora', 'loha', 'lycoris', 'controlnet', 'textual inversion', \n", + " 'embedding', 'ti', 'checkpoint', 'model', 'adapter', 'pony', 'sdxl', 'flux', 'illustrious', 'sd14', 'sd14', 'sd2', 'sd3', 'diffusion', 'stable', 'hunyuan']\n", + " for term in lora_terms:\n", + " name = re.sub(rf'\\b{term}\\b', '', name, flags=re.IGNORECASE)\n", + " \n", + " # Remove content in parentheses or brackets (often metadata)\n", + " name = re.sub(r'\\([^)]*\\)', '', name)\n", + " name = re.sub(r'\\[[^\\]]*\\]', '', name)\n", + " \n", + " # Remove special characters like 「」\n", + " name = re.sub(r'[「」『』【】〈〉《》]', '', name)\n", + " \n", + " # Handle pipe - keep first part\n", + " if '|' in name:\n", + " name = name.split('|')[0]\n", + " \n", + " # Handle forward slash - keep first part\n", + " if '/' in name:\n", + " name = name.split('/')[0]\n", + " \n", + " # Replace underscores with spaces\n", + " name = name.replace('_', ' ')\n", + " \n", + " # Remove multiple spaces\n", + " name = re.sub(r'\\s+', ' ', name)\n", + " \n", + " # Strip\n", + " name = name.strip()\n", + " \n", + " return name\n", + "\n", + "def extract_person_name(text: str) -> str:\n", + " \"\"\"\n", + " Use spaCy NER to extract person names from text.\n", + " Falls back to cleaned text if no PERSON entity found.\n", + " \"\"\"\n", + " if not text:\n", + " return \"\"\n", + " \n", + " # Run spaCy NER\n", + " doc = nlp(text)\n", + " \n", + " # Extract PERSON entities\n", + " person_entities = [ent.text for ent in doc.ents if ent.label_ == \"PERSON\"]\n", + " \n", + " if person_entities:\n", + " # Return the first (usually longest/best) person name\n", + " return person_entities[0].strip()\n", + " \n", + " # If no PERSON entity found, try to extract capitalized words (likely names)\n", + " # This helps with names spaCy might miss\n", + " words = text.split()\n", + " capitalized_words = [w for w in words if w and w[0].isupper() and len(w) > 1]\n", + " \n", + " if capitalized_words:\n", + " # Join first few capitalized words (likely the name)\n", + " return ' '.join(capitalized_words[:3]).strip()\n", + " \n", + " # Last resort: return cleaned text\n", + " return text.strip()\n", + "\n", + "def clean_name_with_spacy(name: str) -> str:\n", + " \"\"\"\n", + " Complete name cleaning pipeline with spaCy NER.\n", + " \n", + " Pipeline:\n", + " 1. Translate leetspeak (4→a, 3→e, 1→i, etc.)\n", + " 2. Remove noise (emoji, version tags, LoRA terms)\n", + " 3. Use spaCy to extract PERSON entities\n", + " 4. Fallback to capitalized words or cleaned text\n", + " \"\"\"\n", + " # Step 1 & 2: Preprocess (leetspeak + noise removal)\n", + " preprocessed = preprocess_for_ner(name)\n", + " \n", + " if not preprocessed:\n", + " return \"\"\n", + " \n", + " # Step 3: Extract person name using spaCy NER\n", + " person_name = extract_person_name(preprocessed)\n", + " \n", + " return person_name\n", + "\n", + "# Apply name cleaning with spaCy\n", + "print(\"\\n🔄 Processing names with spaCy NER...\")\n", + "df['real_name'] = df['name'].apply(clean_name_with_spacy)\n", + "\n", + "# Show examples with detailed comparison\n", + "print(\"\\n📊 Name cleaning examples (with spaCy NER):\")\n", + "print(\"=\" * 100)\n", + "print(f\"{'Original Name':<50} | {'Cleaned Name':<30}\")\n", + "print(\"=\" * 100)\n", + "\n", + "examples = df[['name', 'real_name']].head(30)\n", + "shown = 0\n", + "for idx, row in examples.iterrows():\n", + " if row['name'] != row['real_name'] and shown < 20:\n", + " print(f\"{row['name']:<50} | {row['real_name']:<30}\")\n", + " shown += 1\n", + "\n", + "print(\"=\" * 100)\n", + "\n", + "# Show specific test cases\n", + "print(\"\\n🧪 Leetspeak translation examples:\")\n", + "test_names = ['4kira LoRA', '3mma Watson v2', '1rene LORA', 'L3vi Ackerman']\n", + "for test in test_names:\n", + " result = clean_name_with_spacy(test)\n", + " print(f\" {test:<30} -> {result}\")\n", + "\n", + "# Statistics\n", + "print(f\"\\n📈 Statistics:\")\n", + "print(f\" Total rows: {len(df)}\")\n", + "print(f\" Non-empty names: {(df['real_name'] != '').sum()}\")\n", + "print(f\" Empty names: {(df['real_name'] == '').sum()}\")\n", + "\n", + "# Show some examples of what spaCy identified\n", + "print(\"\\n🎯 Sample spaCy NER results:\")\n", + "sample_names = df['real_name'].head(20).tolist()\n", + "for i, name in enumerate(sample_names[:10], 1):\n", + " if name:\n", + " print(f\" {i}. {name}\")\n", + "\n", + "print(f\"\\n✅ Cleaned {len(df)} names using spaCy NER\")\n", + "\n", + "# Save intermediate result\n", + "output_step1 = current_dir.parent / \"data/CSV/model_adapter/real_person_adapter_step_01_NER.csv\"\n", + "df.to_csv(output_step1, index=False)\n", + "print(f\"💾 Saved to {output_step1}\")\n" + ] + }, + { + "cell_type": "markdown", + "id": "64687c72", + "metadata": {}, + "source": [ + "#### STEP 02: Nationality tag to Country hint\n", + "here tags related to nationality gets converted to the country equivalent." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "d6eaef5b", + "metadata": {}, + "outputs": [], + "source": [ + "import pandas as pd\n", + "from pathlib import Path\n", + "\n", + "# Set up paths\n", + "current_dir = Path.cwd()\n", + "countries_file = current_dir.parent / \"misc/lists/countries.csv\"\n", + "professions_file = current_dir.parent / \"misc/lists/professions.csv\"\n", + "input_file = current_dir.parent / \"data/CSV/model_adapter/real_person_adapter_step_01_NER.csv\"\n", + "output_file = current_dir.parent / \"data/CSV/model_adapter/real_person_adapter_step_02_NER.csv\"\n", + "\n", + "# Load datasets\n", + "poi_df = pd.read_csv(input_file)\n", + "countries_df = pd.read_csv(countries_file)\n", + "professions_df = pd.read_csv(professions_file)\n", + "\n", + "# Define uninhabited or non-relevant territories to exclude\n", + "excluded_territories = {\n", + " 'isle of man', 'bouvet island', 'heard island and mcdonald islands',\n", + " 'french southern territories', 'south georgia and the south sandwich islands',\n", + " 'svalbard and jan mayen', 'british indian ocean territory', 'antarctica',\n", + " 'christmas island', 'cocos (keeling) islands', 'norfolk island',\n", + " 'pitcairn', 'tokelau', 'united states minor outlying islands',\n", + " 'wallis and futuna', 'western sahara'\n", + "}\n", + "\n", + "# Step 1: Combine tags into one lowercase list\n", + "def combine_tags(row):\n", + " return [str(row[f\"tag_{i}\"]).strip().lower() for i in range(1, 8) if pd.notna(row.get(f\"tag_{i}\"))]\n", + "\n", + "poi_df[\"tags\"] = poi_df.apply(combine_tags, axis=1)\n", + "\n", + "# Step 2: Build tag → (country, nationality) mapping with PRIORITIES\n", + "tag_to_country_nationality = {}\n", + "# We'll use a priority score: direct country name = 3, nationality = 2, word parts = 1\n", + "\n", + "for _, row in countries_df.iterrows():\n", + " country = str(row[\"en_short_name\"]).strip()\n", + " nationality = str(row[\"nationality\"]).strip()\n", + " \n", + " # Skip excluded territories\n", + " if country.lower() in excluded_territories:\n", + " continue\n", + "\n", + " country_lc = country.lower()\n", + " nationality_lc = nationality.lower()\n", + "\n", + " # Store as (country, nationality, priority)\n", + " # Exact country name match = highest priority\n", + " if country_lc not in tag_to_country_nationality:\n", + " tag_to_country_nationality[country_lc] = (country, \"\", 3)\n", + " \n", + " # Exact nationality match = medium priority \n", + " if nationality_lc not in tag_to_country_nationality:\n", + " tag_to_country_nationality[nationality_lc] = (\"\", nationality, 2)\n", + " \n", + " # No-space versions\n", + " country_no_space = country_lc.replace(\" \", \"\")\n", + " nationality_no_space = nationality_lc.replace(\" \", \"\")\n", + " \n", + " if country_no_space not in tag_to_country_nationality:\n", + " tag_to_country_nationality[country_no_space] = (country, \"\", 3)\n", + " if nationality_no_space not in tag_to_country_nationality:\n", + " tag_to_country_nationality[nationality_no_space] = (\"\", nationality, 2)\n", + "\n", + " # Word parts = lowest priority (only for longer words to avoid false matches)\n", + " for part in country_lc.split():\n", + " if len(part) > 4: # Only words longer than 4 chars\n", + " if part not in tag_to_country_nationality:\n", + " tag_to_country_nationality[part] = (country, \"\", 1)\n", + " for part in nationality_lc.split():\n", + " if len(part) > 4:\n", + " if part not in tag_to_country_nationality:\n", + " tag_to_country_nationality[part] = (\"\", nationality, 1)\n", + "\n", + "print(f\"Built country/nationality mapping with {len(tag_to_country_nationality)} entries\")\n", + "\n", + "# Step 3: Infer likely_country and likely_nationality by checking ALL tags\n", + "def infer_country_and_nationality(tags):\n", + " \"\"\"\n", + " Check ALL tags and return the best match based on priority.\n", + " Priority: exact country name > nationality > word parts\n", + " \"\"\"\n", + " best_match = None\n", + " best_priority = 0\n", + " \n", + " for tag in tags:\n", + " # Try cleaned version (no spaces)\n", + " cleaned = tag.replace(\" \", \"\").lower()\n", + " \n", + " # Check cleaned version\n", + " if cleaned in tag_to_country_nationality:\n", + " country, nationality, priority = tag_to_country_nationality[cleaned]\n", + " if priority > best_priority and country and country.lower() not in excluded_territories:\n", + " best_match = (country, nationality)\n", + " best_priority = priority\n", + " \n", + " # Also check original tag\n", + " if tag in tag_to_country_nationality:\n", + " country, nationality, priority = tag_to_country_nationality[tag]\n", + " if priority > best_priority and country and country.lower() not in excluded_territories:\n", + " best_match = (country, nationality)\n", + " best_priority = priority\n", + " \n", + " if best_match:\n", + " return pd.Series(best_match)\n", + " return pd.Series([\"\", \"\"])\n", + "\n", + "poi_df[[\"likely_country\", \"likely_nationality\"]] = poi_df[\"tags\"].apply(infer_country_and_nationality)\n", + "\n", + "# Step 4: Build tag → profession mapping\n", + "profession_alias_map = {}\n", + "\n", + "for _, row in professions_df.iterrows():\n", + " canonical = str(row['profession']).strip().lower()\n", + " profession_alias_map[canonical] = canonical\n", + " for alias_col in ['alias_1', 'alias_2', 'alias_3']:\n", + " alias = row.get(alias_col)\n", + " if pd.notna(alias):\n", + " profession_alias_map[str(alias).strip().lower()] = canonical\n", + "\n", + "# Step 5: Infer likely profession from tags\n", + "def infer_profession_from_tags(tags):\n", + " matched = []\n", + " for tag in tags:\n", + " cleaned = tag.strip().lower()\n", + " if cleaned in profession_alias_map:\n", + " matched.append(profession_alias_map[cleaned])\n", + "\n", + " if not matched:\n", + " return \"\"\n", + " if \"celebrity\" in matched and len(set(matched)) > 1:\n", + " # Drop 'celebrity' if other professions are present\n", + " matched = [m for m in matched if m != \"celebrity\"]\n", + "\n", + " return matched[0] # Return the first specific match\n", + "\n", + "\n", + "poi_df[\"likely_profession\"] = poi_df[\"tags\"].apply(infer_profession_from_tags)\n", + "\n", + "# Step 6: Save enriched dataset\n", + "poi_df.to_csv(output_file, index=False)\n", + "\n", + "# Preview results\n", + "print(f\"\\nProcessed {len(poi_df)} rows\")\n", + "print(f\"Rows with country: {(poi_df['likely_country'] != '').sum()}\")\n", + "print(f\"Rows with nationality: {(poi_df['likely_nationality'] != '').sum()}\")\n", + "print(f\"Rows with profession: {(poi_df['likely_profession'] != '').sum()}\")\n", + "\n", + "print(f\"\\nTop 10 countries:\")\n", + "print(poi_df[poi_df['likely_country'] != '']['likely_country'].value_counts().head(10))\n" + ] + }, + { + "cell_type": "markdown", + "id": "4a4a58b3", + "metadata": {}, + "source": [ + "## LLM ANNOTATION" + ] + }, + { + "cell_type": "markdown", + "id": "b298844d", + "metadata": {}, + "source": [ + "#### Model Configurations" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "39f3d65e", + "metadata": {}, + "outputs": [], + "source": [ + "import pandas as pd\n", + "import json\n", + "import time\n", + "import re\n", + "from pathlib import Path\n", + "from tqdm import tqdm\n", + "import torch\n", + "from transformers import AutoModelForCausalLM, AutoTokenizer, BitsAndBytesConfig\n", + "import signal\n", + "from contextlib import contextmanager\n", + "\n", + "# Configuration\n", + "current_dir = Path.cwd()\n", + "CACHE_DIR = current_dir.parent / \"data/models\"\n", + "CACHE_DIR.mkdir(parents=True, exist_ok=True)\n", + "\n", + "# Model configurations\n", + "MODEL_CONFIGS = {\n", + " 'mistral': {\n", + " 'name': 'mistralai/Mistral-7B-Instruct-v0.3',\n", + " 'dtype': torch.bfloat16,\n", + " 'quantization': None,\n", + " 'generation_params': {\n", + " 'max_new_tokens': 512,\n", + " 'temperature': 0.05,\n", + " 'do_sample': True,\n", + " 'top_p': 0.8,\n", + " }\n", + " },\n", + " 'gemma': {\n", + " 'name': 'google/gemma-3-27b-it',\n", + " 'dtype': torch.bfloat16,\n", + " 'quantization': None,\n", + " 'generation_params': {\n", + " 'max_new_tokens': 512,\n", + " 'temperature': 0.1,\n", + " 'do_sample': True,\n", + " 'top_p': 0.9,\n", + " }\n", + " },\n", + " 'qwen': {\n", + " 'name': 'Qwen/Qwen2.5-32B-Instruct',\n", + " 'dtype': None, # Will use quantization\n", + " 'quantization': BitsAndBytesConfig(\n", + " load_in_8bit=True,\n", + " llm_int8_threshold=6.0,\n", + " llm_int8_has_fp16_weight=False\n", + " ),\n", + " 'generation_params': {\n", + " 'max_new_tokens': 100,\n", + " 'temperature': 0.1,\n", + " 'do_sample': False,\n", + " }\n", + " }\n", + "}\n", + "\n", + "PROFESSION_CATEGORIES = [\n", + " \"actor\",\n", + " \"adult performer\",\n", + " \"singer/musician\",\n", + " \"model\",\n", + " \"online personality\",\n", + " \"public figure\",\n", + " \"voice actor/ASMR\",\n", + " \"sports professional\",\n", + " \"tv personality\"\n", + "]\n" + ] + }, + { + "cell_type": "markdown", + "id": "c215b38c", + "metadata": {}, + "source": [ + "#### Load Model Function" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "cfb5b13e", + "metadata": {}, + "outputs": [], + "source": [ + "def load_model(model_type='mistral'):\n", + " \"\"\"\n", + " Load model and tokenizer based on type.\n", + " \n", + " Args:\n", + " model_type: 'mistral', 'gemma', or 'qwen'\n", + " \n", + " Returns:\n", + " tuple: (model, tokenizer, config)\n", + " \"\"\"\n", + " if model_type not in MODEL_CONFIGS:\n", + " raise ValueError(f\"Unknown model type: {model_type}. Choose from {list(MODEL_CONFIGS.keys())}\")\n", + " \n", + " config = MODEL_CONFIGS[model_type]\n", + " model_name = config['name']\n", + " \n", + " device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n", + " print(f\"Loading model: {model_name}\")\n", + " print(f\"Cache directory: {CACHE_DIR}\")\n", + " print(f\"Device: {device}\\n\")\n", + " \n", + " if device == \"cpu\":\n", + " print(\"⚠️ WARNING: No GPU detected! Inference will be VERY slow.\")\n", + " \n", + " # Load tokenizer\n", + " try:\n", + " tokenizer = AutoTokenizer.from_pretrained(\n", + " model_name,\n", + " cache_dir=str(CACHE_DIR),\n", + " use_fast=True\n", + " )\n", + " except:\n", + " tokenizer = AutoTokenizer.from_pretrained(\n", + " model_name,\n", + " cache_dir=str(CACHE_DIR),\n", + " use_fast=False\n", + " )\n", + " \n", + " if tokenizer.pad_token is None:\n", + " tokenizer.pad_token = tokenizer.eos_token\n", + " \n", + " # Load model\n", + " model_kwargs = {\n", + " 'cache_dir': str(CACHE_DIR),\n", + " 'device_map': 'auto',\n", + " 'trust_remote_code': False\n", + " }\n", + " \n", + " if config['quantization']:\n", + " model_kwargs['quantization_config'] = config['quantization']\n", + " else:\n", + " model_kwargs['torch_dtype'] = config['dtype']\n", + " \n", + " model = AutoModelForCausalLM.from_pretrained(model_name, **model_kwargs)\n", + " model.eval()\n", + " \n", + " # Check VRAM\n", + " if torch.cuda.is_available():\n", + " vram_gb = torch.cuda.max_memory_allocated() / 1024**3\n", + " print(f\"VRAM used: {vram_gb:.2f} GB\\n\")\n", + " \n", + " return model, tokenizer, config\n" + ] + }, + { + "cell_type": "markdown", + "id": "11b2221a", + "metadata": {}, + "source": [ + "#### Inference Code" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "229f96bd", + "metadata": {}, + "outputs": [], + "source": [ + "@contextmanager\n", + "def timeout(duration):\n", + " \"\"\"Context manager for timeout.\"\"\"\n", + " def handler(signum, frame):\n", + " raise TimeoutError(\"Operation timed out\")\n", + " \n", + " signal.signal(signal.SIGALRM, handler)\n", + " signal.alarm(duration)\n", + " try:\n", + " yield\n", + " finally:\n", + " signal.alarm(0)\n", + "\n", + "def query_model(prompt, model, tokenizer, config, use_timeout=False):\n", + " \"\"\"\n", + " Query model with given prompt.\n", + " \n", + " Args:\n", + " prompt: Input prompt string\n", + " model: Loaded model\n", + " tokenizer: Loaded tokenizer\n", + " config: Model configuration dict\n", + " use_timeout: Whether to use 60s timeout (for Qwen)\n", + " \n", + " Returns:\n", + " str: Model response or None on error\n", + " \"\"\"\n", + " try:\n", + " device = next(model.parameters()).device\n", + " \n", + " # Format as chat message\n", + " messages = [\n", + " {\"role\": \"system\", \"content\": \"You are a data extraction assistant. Respond with exactly 5 numbered lines containing ONLY values. No labels, no explanations, no prefixes. Follow the format precisely.\"},\n", + " {\"role\": \"user\", \"content\": prompt}\n", + " ]\n", + " \n", + " # Tokenize\n", + " if hasattr(tokenizer, 'apply_chat_template'):\n", + " text = tokenizer.apply_chat_template(\n", + " messages,\n", + " tokenize=False,\n", + " add_generation_prompt=True\n", + " )\n", + " else:\n", + " text = f\"[INST] {prompt} [/INST]\"\n", + " \n", + " inputs = tokenizer([text], return_tensors=\"pt\", padding=True).to(device)\n", + " \n", + " # Generation parameters\n", + " gen_kwargs = config['generation_params'].copy()\n", + " gen_kwargs['pad_token_id'] = tokenizer.eos_token_id\n", + " \n", + " # Generate\n", + " generation_fn = lambda: model.generate(**inputs, **gen_kwargs)\n", + " \n", + " if use_timeout:\n", + " with timeout(60):\n", + " with torch.no_grad():\n", + " outputs = generation_fn()\n", + " else:\n", + " with torch.no_grad():\n", + " outputs = generation_fn()\n", + " \n", + " # Decode\n", + " generated_ids = outputs[0][inputs['input_ids'].shape[1]:]\n", + " response = tokenizer.decode(generated_ids, skip_special_tokens=True)\n", + " \n", + " return response.strip()\n", + " \n", + " except TimeoutError:\n", + " print(f\"[ERROR] Generation timed out after 60 seconds\")\n", + " return None\n", + " except Exception as e:\n", + " print(f\"[ERROR] Generation failed: {e}\")\n", + " return None\n" + ] + }, + { + "cell_type": "markdown", + "id": "88f005f8", + "metadata": {}, + "source": [ + "#### Prompt creation" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "dfe05463", + "metadata": {}, + "outputs": [], + "source": [ + "def create_prompt(row):\n", + " \"\"\"Create annotation prompt from row data.\"\"\"\n", + " name = row['real_name'] if pd.notna(row.get('real_name')) else row.get('name', '')\n", + " \n", + " # Gather hints\n", + " hints = []\n", + " if pd.notna(row.get('likely_profession')):\n", + " hints.append(str(row['likely_profession']))\n", + " if pd.notna(row.get('likely_nationality')):\n", + " hints.append(str(row['likely_nationality']))\n", + " if pd.notna(row.get('likely_country')):\n", + " hints.append(str(row['likely_country']))\n", + " \n", + " # Add tags if needed\n", + " if len(hints) < 3:\n", + " for i in range(1, 8):\n", + " tag_col = f'tag_{i}'\n", + " if tag_col in row and pd.notna(row[tag_col]):\n", + " tag_val = str(row[tag_col])\n", + " if tag_val not in hints:\n", + " hints.append(tag_val)\n", + " if len(hints) >= 5:\n", + " break\n", + " \n", + " hint_text = \", \".join(hints[:5]) if hints else \"none\"\n", + " \n", + " return f\"\"\"Extract information about '{name}' ({hint_text}).\n", + "\n", + "Respond with EXACTLY 5 numbered lines. Each line must contain ONLY the value, no labels or extra text.\n", + "\n", + "FORMAT REQUIREMENTS:\n", + "1. Full legal name in Western order (first last). VALUE ONLY.\n", + "2. Stage names/aliases, comma-separated. If none, write \"None\". VALUE ONLY.\n", + "3. Gender: MUST be exactly one word: Male, Female, Other, or Unknown. VALUE ONLY.\n", + "4. Professions: Choose up to 3 from this list ONLY: actor, adult performer, singer/musician, model, online personality, public figure, voice actor/ASMR, sports professional, tv personality. Comma-separated. VALUE ONLY.\n", + "5. Primary country: Country name only (e.g., \"China\", \"United States\", \"Colombia\"). VALUE ONLY.\n", + "\n", + "RULES:\n", + "- Professions MUST match the exact categories listed (actress = actor)\n", + "- \"online personality\" includes streamers, cosplayers, YouTubers, influencers\n", + "- \"public figure\" includes politicians, activists, journalists, authors\n", + "- Use \"Unknown\" when uncertain or for fictional characters\n", + "- NO explanatory text, NO labels like \"Gender:\", NO prefixes\n", + "- For multi-role people, list up to 3 categories by relevance\n", + "\n", + "EXAMPLE FORMAT:\n", + "1. Taylor Swift\n", + "2. None\n", + "3. Female\n", + "4. singer/musician, public figure\n", + "5. United States\"\"\"\n" + ] + }, + { + "cell_type": "markdown", + "id": "854fa668", + "metadata": {}, + "source": [ + "#### Response parsing code" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "1a4be2ee", + "metadata": {}, + "outputs": [], + "source": [ + "def parse_response(response):\n", + " \"\"\"Parse model response into structured fields.\"\"\"\n", + " if not response:\n", + " return {\n", + " 'full_name': 'Unknown',\n", + " 'aliases': 'Unknown',\n", + " 'gender': 'Unknown',\n", + " 'profession_llm': 'Unknown',\n", + " 'country': 'Unknown'\n", + " }\n", + " \n", + " lines = [line.strip() for line in response.split('\\n') if line.strip()]\n", + " \n", + " fields = {\n", + " 'full_name': 'Unknown',\n", + " 'aliases': 'Unknown',\n", + " 'gender': 'Unknown',\n", + " 'profession_llm': 'Unknown',\n", + " 'country': 'Unknown'\n", + " }\n", + " \n", + " for line in lines:\n", + " if line.startswith('1.'):\n", + " fields['full_name'] = line[2:].strip()\n", + " elif line.startswith('2.'):\n", + " fields['aliases'] = line[2:].strip()\n", + " elif line.startswith('3.'):\n", + " gender_raw = line[2:].strip()\n", + " gender_raw = re.sub(r'^(Gender:|gender:)\\s*', '', gender_raw, flags=re.IGNORECASE)\n", + " gender_match = re.search(r'\\b(Male|Female|Other|Unknown)\\b', gender_raw, re.IGNORECASE)\n", + " fields['gender'] = gender_match.group(1).capitalize() if gender_match else gender_raw\n", + " elif line.startswith('4.'):\n", + " fields['profession_llm'] = line[2:].strip()\n", + " elif line.startswith('5.'):\n", + " country_raw = line[2:].strip()\n", + " country_raw = re.sub(r'^(Primary\\s+)?(associated\\s+)?country:\\s*', '', country_raw, flags=re.IGNORECASE)\n", + " fields['country'] = country_raw\n", + " \n", + " return fields\n" + ] + }, + { + "cell_type": "markdown", + "id": "7e2f7a86", + "metadata": {}, + "source": [ + "#### CSV annotation" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "5f3dd5d6", + "metadata": {}, + "outputs": [], + "source": [ + "def annotate_dataset(model_type='mistral', test_mode=False, test_size=100, max_rows=50862, save_interval=10):\n", + " \"\"\"\n", + " Annotate dataset using specified model.\n", + " \n", + " Args:\n", + " model_type: 'mistral', 'gemma', or 'qwen'\n", + " test_mode: If True, only process test_size rows\n", + " test_size: Number of rows to process in test mode\n", + " max_rows: Maximum rows to process\n", + " save_interval: Save progress every N rows\n", + " \"\"\"\n", + " # Setup paths\n", + " input_file = current_dir.parent / \"data/CSV/model_adapter/real_person_adapter_step_02_NER.csv\"\n", + " output_file = current_dir.parent / f\"data/CSV/{model_type}_local_annotated_POI{'_test' if test_mode else ''}.csv\"\n", + " index_file = current_dir.parent / f\"misc/query_indicies/{model_type}_local_query_index.txt\"\n", + " index_file.parent.mkdir(parents=True, exist_ok=True)\n", + " \n", + " # Load model\n", + " model, tokenizer, config = load_model(model_type)\n", + " \n", + " # Load data\n", + " print(f\"Loaded {len(df)} rows from input file\")\n", + " df = pd.read_csv(input_file)\n", + " \n", + " # Merge existing annotations if available\n", + " if output_file.exists():\n", + " existing_df = pd.read_csv(output_file)\n", + " annotation_cols = ['full_name', 'aliases', 'gender', 'profession_llm', 'country']\n", + " for col in annotation_cols:\n", + " if col in existing_df.columns:\n", + " df[col] = existing_df[col][:len(df)]\n", + " \n", + " # Apply limits\n", + " if test_mode:\n", + " df = df.head(test_size).copy()\n", + " elif max_rows:\n", + " df = df.head(max_rows).copy()\n", + " \n", + " # Create prompts\n", + " df['prompt'] = df.apply(create_prompt, axis=1)\n", + " \n", + " # Load progress index\n", + " current_index = 0\n", + " if index_file.exists():\n", + " try:\n", + " current_index = int(index_file.read_text().strip())\n", + " except:\n", + " current_index = 0\n", + " \n", + " print(f\"Resuming from index {current_index}\")\n", + " \n", + " # Process rows\n", + " use_timeout = (model_type == 'qwen')\n", + " \n", + " for i in tqdm(range(current_index, len(df)), desc=f\"{model_type.capitalize()} Annotation\"):\n", + " prompt = df.at[i, \"prompt\"]\n", + " \n", + " # Query with retries\n", + " response = None\n", + " for attempt in range(3):\n", + " response = query_model(prompt, model, tokenizer, config, use_timeout)\n", + " \n", + " if response and len(response.strip()) > 10:\n", + " break\n", + " \n", + " print(f\"⚠️ Row {i}: Empty or invalid response, retry {attempt+1}/3\")\n", + " time.sleep(0.5)\n", + " \n", + " # Skip if invalid\n", + " if not response or len(response.strip()) <= 10:\n", + " print(f\"❌ Row {i}: failed after retries, skipping\")\n", + " continue\n", + " \n", + " # Parse and validate\n", + " parsed = parse_response(response)\n", + " \n", + " if all(v == \"Unknown\" for v in parsed.values()):\n", + " print(f\"❌ Row {i}: parsed as all Unknown, skipping\")\n", + " continue\n", + " \n", + " # Write fields\n", + " for key, value in parsed.items():\n", + " df.at[i, key] = value\n", + " \n", + " current_index = i + 1\n", + " \n", + " # GPU cleanup\n", + " if torch.cuda.is_available():\n", + " torch.cuda.empty_cache()\n", + " torch.cuda.synchronize()\n", + " \n", + " # Save progress\n", + " if (i + 1) % save_interval == 0 or (i + 1) == len(df):\n", + " df.to_csv(output_file, index=False)\n", + " index_file.write_text(str(current_index))\n", + " print(f\"💾 Progress saved after row {i+1}\")\n", + " \n", + " # Final save\n", + " df.to_csv(output_file, index=False)\n", + " index_file.write_text(str(current_index))\n", + " print(f\"✓ Finished annotation with {model_type}\")\n" + ] + }, + { + "cell_type": "markdown", + "id": "55da2f4c", + "metadata": {}, + "source": [ + "### Usage Examples\n", + "Run annotation with your chosen model." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "351ea40c", + "metadata": {}, + "outputs": [], + "source": [ + "# Example 1: Annotate with Mistral (13.5 GB VRAM)\n", + "# annotate_dataset(model_type='mistral', test_mode=False)\n", + "\n", + "# Example 2: Annotate with Gemma (56.3 GB VRAM)\n", + "# annotate_dataset(model_type='gemma', test_mode=False)\n", + "\n", + "# Example 3: Annotate with Qwen (32.7 GB VRAM, 8-bit)\n", + "# annotate_dataset(model_type='qwen', test_mode=False)\n", + "\n", + "# Test mode (first 100 rows)\n", + "# annotate_dataset(model_type='mistral', test_mode=True, test_size=100)\n" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "e8203abc-e7c3-4cb6-aaeb-fdc6933981fc", + "metadata": {}, + "outputs": [], + "source": [ + "import pandas as pd\n", + "import json\n", + "import time\n", + "import re\n", + "from pathlib import Path\n", + "from tqdm import tqdm\n", + "import torch\n", + "from transformers import AutoModelForCausalLM, AutoTokenizer, BitsAndBytesConfig\n", + "import signal\n", + "from contextlib import contextmanager\n", + "\n", + "current_dir = Path.cwd()\n", + "input_file = current_dir.parent / \"data/CSV/model_adapter/real_person_adapter_step_02_NER.csv\"\n", + "professions_file = current_dir.parent / \"misc/lists/professions.csv\"\n", + "professions_mapped_file = current_dir.parent / \"misc/lists/professions_mapped.csv\"\n", + "# === PROCESS DATA ===\n", + "\n", + "\n", + "# === CONFIGURATION ===\n", + "TEST_MODE = False\n", + "TEST_SIZE = 100\n", + "MAX_ROWS = 50862\n", + "SAVE_INTERVAL = 10\n", + "\n", + "\n", + "index_file = current_dir.parent / \"misc/query_indicies/eurollm_local_query_index.txt\"\n", + "output_file = current_dir.parent / f\"data/CSV/eurollm_local_annotated_POI{'_test' if TEST_MODE else ''}.csv\"\n", + "\n", + "# Model settings\n", + "MODEL_NAME = \"utter-project/EuroLLM-9B\"\n", + "#MODEL_NAME = \"Qwen/Qwen2.5-32B-Instruct\"\n", + "#MODEL_NAME = \"Qwen/Qwen2.5-14B-Instruct\"\n", + "#MODEL_NAME = \"Qwen/Qwen3-235B-A22B-Instruct-2507-FP8\"\n", + "#MODEL_NAME = \"mistralai/Mistral-7B-Instruct-v0.3\"\n", + "CACHE_DIR = current_dir.parent / \"data/models\"\n", + "CACHE_DIR.mkdir(parents=True, exist_ok=True)\n", + "\n", + "# Define the SPECIFIC profession categories\n", + "PROFESSION_CATEGORIES = [\n", + " \"actor\",\n", + " \"adult performer\",\n", + " \"singer/musician\",\n", + " \"model\",\n", + " \"online personality\",\n", + " \"public figure\",\n", + " \"voice actor/ASMR\",\n", + " \"sports professional\",\n", + " \"tv personality\"\n", + "]\n", + "\n", + "# === LOAD MODEL ===\n", + "print(f\"Loading model: {MODEL_NAME}\")\n", + "print(f\"Cache directory: {CACHE_DIR}\")\n", + "print(f\"This may take a while on first run...\\n\")\n", + "\n", + "# Check GPU availability\n", + "device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n", + "print(f\"Device: {device}\")\n", + "\n", + "if device == \"cpu\":\n", + " print(\"⚠️ WARNING: No GPU detected! Inference will be VERY slow.\")\n", + " print(\" Consider using a GPU or reducing model size.\")\n", + "\n", + "# Get HF token from credentials file\n", + "import os\n", + "credentials_dir = current_dir.parent / \"misc/credentials\"\n", + "hf_token_file = credentials_dir / \"hf_token.txt\"\n", + "\n", + "HF_TOKEN = None\n", + "if hf_token_file.exists():\n", + " HF_TOKEN = hf_token_file.read_text().strip()\n", + " print(\"✅ HF token loaded from credentials file\")\n", + "else:\n", + " print(\"⚠️ HF token file not found at:\", hf_token_file)\n", + " print(\" The script will try to use cached credentials from 'huggingface-cli login'\")\n", + " print(\" Or create the file: misc/credentials/hf_token.txt with your token\")\n", + " HF_TOKEN = None # Will use cached token if available\n", + "\n", + "# Load tokenizer\n", + "print(\"Loading tokenizer...\")\n", + "try:\n", + " tokenizer = AutoTokenizer.from_pretrained(\n", + " MODEL_NAME,\n", + " cache_dir=str(CACHE_DIR),\n", + " use_fast=True,\n", + " token=HF_TOKEN\n", + " )\n", + "except Exception as e:\n", + " print(f\"Failed with use_fast=True, trying use_fast=False...\")\n", + " tokenizer = AutoTokenizer.from_pretrained(\n", + " MODEL_NAME,\n", + " cache_dir=str(CACHE_DIR),\n", + " use_fast=False,\n", + " token=HF_TOKEN\n", + " )\n", + "\n", + "# Ensure pad token is set\n", + "if tokenizer.pad_token is None:\n", + " tokenizer.pad_token = tokenizer.eos_token\n", + "\n", + "print(\"✅ Tokenizer loaded\")\n", + "\n", + "# Configure 8-bit quantization for A100\n", + "print(\"Configuring 8-bit quantization...\")\n", + "quantization_config = BitsAndBytesConfig(\n", + " load_in_8bit=True,\n", + " llm_int8_threshold=6.0,\n", + " llm_int8_has_fp16_weight=False\n", + ")\n", + "\n", + "# Load model with 8-bit quantization\n", + "print(\"Loading model with 8-bit quantization (this may take several minutes)...\")\n", + "model = AutoModelForCausalLM.from_pretrained(\n", + " MODEL_NAME,\n", + " cache_dir=str(CACHE_DIR),\n", + " quantization_config=quantization_config,\n", + " device_map=\"auto\",\n", + " trust_remote_code=False,\n", + " token=HF_TOKEN\n", + ")\n", + "model.eval()\n", + "print(\"✅ Model loaded with 8-bit quantization\")\n", + "\n", + "# Check VRAM usage\n", + "if torch.cuda.is_available():\n", + " vram_gb = torch.cuda.max_memory_allocated() / 1024**3\n", + " print(f\"VRAM used: {vram_gb:.2f} GB\\n\")\n", + "\n", + "# === LOAD DATA ===\n", + "if output_file.exists():\n", + " print(\"Loading annotated CSV...\")\n", + " df = pd.read_csv(output_file)\n", + "else:\n", + " print(\"Loading raw input CSV...\")\n", + " df = pd.read_csv(input_file)\n", + "\n", + "\n", + "# Try to load profession mapping files\n", + "try:\n", + " professions_df = pd.read_csv(professions_file)\n", + " print(f\"✅ Loaded professions.csv\")\n", + "except:\n", + " print(\"⚠️ Warning: professions.csv not found\")\n", + "\n", + "try:\n", + " prof_mapped_df = pd.read_csv(professions_mapped_file)\n", + " print(f\"✅ Loaded profession mapping with {len(prof_mapped_df)} categories\")\n", + "except:\n", + " print(\"⚠️ Warning: professions_mapped.csv not found, using default categories\")\n", + "\n", + "profession_str = \", \".join(PROFESSION_CATEGORIES)\n", + "\n", + "print(f\"Loaded {len(df)} rows\")\n", + "print(f\"\\nProfession categories ({len(PROFESSION_CATEGORIES)}):\")\n", + "for cat in PROFESSION_CATEGORIES:\n", + " print(f\" - {cat}\")\n", + "\n", + "if TEST_MODE:\n", + " print(f\"\\nRunning in TEST MODE with {TEST_SIZE} samples\")\n", + " df = df.head(TEST_SIZE).copy()\n", + "elif MAX_ROWS:\n", + " df = df.head(MAX_ROWS).copy()\n", + "\n", + "# === CREATE PROMPTS (OPTIMIZED FOR CLEAN OUTPUTS) ===\n", + "def create_prompt(row):\n", + " \"\"\"Create prompt for EuroLLM annotation with strict formatting requirements.\"\"\"\n", + " name = row['real_name'] if pd.notna(row.get('real_name')) else row.get('name', '')\n", + " \n", + " # Gather hints\n", + " hints = []\n", + " if pd.notna(row.get('likely_profession')):\n", + " hints.append(str(row['likely_profession']))\n", + " if pd.notna(row.get('likely_nationality')):\n", + " hints.append(str(row['likely_nationality']))\n", + " if pd.notna(row.get('likely_country')):\n", + " hints.append(str(row['likely_country']))\n", + " \n", + " # Add tags if we don't have enough hints\n", + " if len(hints) < 3:\n", + " for i in range(1, 8):\n", + " tag_col = f'tag_{i}'\n", + " if tag_col in row and pd.notna(row[tag_col]):\n", + " tag_val = str(row[tag_col])\n", + " if tag_val not in hints:\n", + " hints.append(tag_val)\n", + " if len(hints) >= 5:\n", + " break\n", + " \n", + " hint_text = \", \".join(hints[:5]) if hints else \"none\"\n", + " \n", + " return f\"\"\"Extract information about '{name}'. \n", + "Context hints (DO NOT copy these as professions): {hint_text}\n", + "\n", + "Respond with EXACTLY 5 numbered lines. Each line must contain ONLY the value, no labels or extra text.\n", + "\n", + "FORMAT REQUIREMENTS:\n", + "1. Full legal name in Western order (first last). VALUE ONLY.\n", + "2. Stage names/aliases, comma-separated. If none, write \"None\". VALUE ONLY.\n", + "3. Gender: MUST be exactly one word: Male, Female, Other, or Unknown. VALUE ONLY.\n", + "4. Professions: Choose up to 3 from this list ONLY: actor, adult performer, singer/musician, model, online personality, public figure, voice actor/ASMR, sports professional, tv personality. Comma-separated. VALUE ONLY.\n", + "5. Primary country: Country name only (e.g., \"China\", \"United States\", \"Colombia\"). VALUE ONLY.\n", + "6. If uncertain about an item, write \"Unknown\"\n", + "\n", + "CRITICAL RULES FOR PROFESSIONS (Line 4):\n", + "- ONLY use the exact profession categories listed above\n", + "- DO NOT use descriptive words like \"sexy\", \"photorealistic\", \"celebrity\"\n", + "- DO NOT copy the hint words as professions\n", + "- If uncertain write \"Unknown\"\n", + "- Valid professions are ONLY: actor, adult performer, singer/musician, model, online personality, public figure, voice actor/ASMR, sports professional, tv personality\n", + "- Actress = actor, streamer = online personality, YouTuber = online personality\n", + "\n", + "OTHER RULES:\n", + "- Use \"Unknown\" when uncertain or for fictional characters\n", + "- NO explanatory text, NO labels like \"Gender:\", NO prefixes\n", + "- For multi-role people, list up to 3 categories by relevance\n", + "\n", + "EXAMPLE FORMAT:\n", + "1. Taylor Swift\n", + "2. None\n", + "3. Female\n", + "4. singer/musician, public figure\n", + "5. United States\"\"\"\n", + "\n", + "# Create prompts\n", + "print(\"\\nCreating prompts...\")\n", + "df['prompt'] = df.apply(create_prompt, axis=1)\n", + "print(\"✅ Prompts created\")\n", + "\n", + "@contextmanager\n", + "def timeout(duration):\n", + " def handler(signum, frame):\n", + " raise TimeoutError(\"Operation timed out\")\n", + " \n", + " # Set the signal handler and alarm\n", + " signal.signal(signal.SIGALRM, handler)\n", + " signal.alarm(duration)\n", + " try:\n", + " yield\n", + " finally:\n", + " signal.alarm(0) # Disable the alarm\n", + "\n", + "\n", + "def query_eurollm_local(prompt: str) -> str:\n", + " \"\"\"Query EuroLLM locally via transformers with very low temperature.\"\"\"\n", + " try:\n", + " # Format as chat message for EuroLLM with strict system prompt\n", + " messages = [\n", + " {\"role\": \"system\", \"content\": \"You are a data extraction assistant. Respond with exactly 5 numbered lines containing ONLY values. No labels, no explanations, no prefixes. Follow the format precisely.\"},\n", + " {\"role\": \"user\", \"content\": prompt}\n", + " ]\n", + " \n", + " # Tokenize\n", + " if hasattr(tokenizer, 'apply_chat_template') and tokenizer.chat_template is not None:\n", + " text = tokenizer.apply_chat_template(\n", + " messages,\n", + " tokenize=False,\n", + " add_generation_prompt=True\n", + " )\n", + " else:\n", + " # Fallback for models without chat template\n", + " text = f\"[INST] {prompt} [/INST]\"\n", + " \n", + " inputs = tokenizer([text], return_tensors=\"pt\", padding=True).to(device)\n", + " \n", + " # Generate with timeout and very low temperature\n", + " with timeout(60):\n", + " with torch.no_grad():\n", + " outputs = model.generate(\n", + " **inputs,\n", + " max_new_tokens=100,\n", + " temperature=0.01, # Very low temperature for more deterministic outputs\n", + " do_sample=True, # Must be True when temperature is set\n", + " pad_token_id=tokenizer.eos_token_id\n", + " )\n", + " \n", + " # Decode\n", + " generated_ids = outputs[0][inputs['input_ids'].shape[1]:]\n", + " response = tokenizer.decode(generated_ids, skip_special_tokens=True)\n", + " \n", + " return response.strip()\n", + " \n", + " except TimeoutError:\n", + " print(f\"[ERROR] Generation timed out after 60 seconds\")\n", + " return None\n", + " except Exception as e:\n", + " print(f\"Generation error: {e}\")\n", + " import traceback\n", + " traceback.print_exc()\n", + " return None\n", + "\n", + " \n", + "# === PARSE RESPONSE WITH CLEANING ===\n", + "def parse_response(response):\n", + " \"\"\"Parse EuroLLM response into structured fields with cleaning.\"\"\"\n", + " if not response:\n", + " return {\n", + " 'full_name': 'Unknown',\n", + " 'aliases': 'Unknown',\n", + " 'gender': 'Unknown',\n", + " 'profession_llm': 'Unknown',\n", + " 'country': 'Unknown'\n", + " }\n", + " \n", + " # Valid profession categories\n", + " VALID_PROFESSIONS = {\n", + " \"actor\", \"adult performer\", \"singer/musician\", \"model\", \n", + " \"online personality\", \"public figure\", \"voice actor/asmr\", \n", + " \"sports professional\", \"tv personality\"\n", + " }\n", + " \n", + " # Split into lines and clean\n", + " lines = [line.strip() for line in response.split('\\n') if line.strip()]\n", + " \n", + " # Initialize with Unknown values\n", + " fields = {\n", + " 'full_name': 'Unknown',\n", + " 'aliases': 'Unknown',\n", + " 'gender': 'Unknown',\n", + " 'profession_llm': 'Unknown',\n", + " 'country': 'Unknown'\n", + " }\n", + " \n", + " # Extract information from each numbered line\n", + " for line in lines:\n", + " if line.startswith('1.'):\n", + " fields['full_name'] = line[2:].strip()\n", + " elif line.startswith('2.'):\n", + " fields['aliases'] = line[2:].strip()\n", + " elif line.startswith('3.'):\n", + " # Clean gender field - remove any labels\n", + " gender_raw = line[2:].strip()\n", + " # Remove common prefixes\n", + " gender_raw = re.sub(r'^(Gender:|gender:)\\s*', '', gender_raw, flags=re.IGNORECASE)\n", + " # Extract just the gender word\n", + " gender_match = re.search(r'\\b(Male|Female|Other|Unknown)\\b', gender_raw, re.IGNORECASE)\n", + " fields['gender'] = gender_match.group(1).capitalize() if gender_match else gender_raw\n", + " elif line.startswith('4.'):\n", + " # Clean and validate profession field\n", + " profession_raw = line[2:].strip()\n", + " \n", + " # Split by comma and validate each profession\n", + " professions = [p.strip().lower() for p in profession_raw.split(',')]\n", + " valid_profs = []\n", + " \n", + " for prof in professions:\n", + " # Check if it's a valid profession\n", + " if prof in VALID_PROFESSIONS:\n", + " valid_profs.append(prof)\n", + " # Check for common invalid entries\n", + " elif prof in ['unknown', '']:\n", + " continue\n", + " # Reject descriptive words that aren't professions\n", + " elif prof in ['sexy', 'photorealistic', 'celebrity', 'famous', 'popular', \n", + " 'beautiful', 'attractive', 'hot', 'gorgeous']:\n", + " continue\n", + " # If it looks like it might be close to a valid profession, keep it\n", + " elif any(valid in prof for valid in VALID_PROFESSIONS):\n", + " # Try to extract the valid part\n", + " for valid in VALID_PROFESSIONS:\n", + " if valid in prof:\n", + " valid_profs.append(valid)\n", + " break\n", + " \n", + " # Set the cleaned professions or Unknown if none are valid\n", + " if valid_profs:\n", + " fields['profession_llm'] = ', '.join(valid_profs)\n", + " else:\n", + " fields['profession_llm'] = 'Unknown'\n", + " \n", + " elif line.startswith('5.'):\n", + " # Clean country field - remove any labels\n", + " country_raw = line[2:].strip()\n", + " # Remove common prefixes like \"Primary country:\", \"Country:\", etc.\n", + " country_raw = re.sub(r'^(Primary\\s+)?(associated\\s+)?country:\\s*', '', country_raw, flags=re.IGNORECASE)\n", + " fields['country'] = country_raw\n", + " \n", + " return fields\n", + "\n", + "# === PROCESS DATA ===\n", + "index_file.parent.mkdir(parents=True, exist_ok=True)\n", + "\n", + "# Load index\n", + "current_index = 0\n", + "if index_file.exists():\n", + " try:\n", + " current_index = int(index_file.read_text().strip())\n", + " except:\n", + " current_index = 0\n", + "\n", + "print(f\"Resuming from index {current_index}\")\n", + "\n", + "start_time = time.time()\n", + "\n", + "for i in tqdm(range(current_index, len(df)), desc=\"EuroLLM Local\"):\n", + "\n", + " prompt = df.at[i, \"prompt\"]\n", + "\n", + " # -------- MODEL QUERY WITH RETRIES --------\n", + " response = None\n", + " for attempt in range(3):\n", + " response = query_eurollm_local(prompt)\n", + " \n", + " # DEBUG: Print first few responses to see what's happening\n", + " if i < 5:\n", + " print(f\"\\n=== DEBUG Row {i}, Attempt {attempt+1} ===\")\n", + " print(f\"Response length: {len(response) if response else 0}\")\n", + " print(f\"Response: {response[:500] if response else 'None'}\")\n", + " print(\"=\" * 50)\n", + " \n", + " # Valid response?\n", + " if response and len(response.strip()) > 10:\n", + " break\n", + " \n", + " print(f\"⚠️ Row {i}: Empty or invalid response, retry {attempt+1}/3\")\n", + " time.sleep(0.5)\n", + "\n", + " # If still invalid → DO NOT overwrite previous data\n", + " if not response or len(response.strip()) <= 10:\n", + " print(f\"❌ Row {i}: failed after retries, not writing, not advancing index\")\n", + " continue\n", + "\n", + " parsed = parse_response(response)\n", + "\n", + " # DEBUG: Print first few parsed results\n", + " if i < 5:\n", + " print(f\"\\n=== PARSED Row {i} ===\")\n", + " for key, value in parsed.items():\n", + " print(f\" {key}: {value}\")\n", + " print(\"=\" * 50)\n", + "\n", + " # Additional safety: skip rows that parsed as all 'Unknown'\n", + " if all(v == \"Unknown\" for v in parsed.values()):\n", + " print(f\"❌ Row {i}: parsed as all Unknown (likely model crash); skipping.\")\n", + " continue\n", + "\n", + " # -------- WRITE PARSED FIELDS SAFELY --------\n", + " for key, value in parsed.items():\n", + " df.at[i, key] = value\n", + "\n", + " # Advance progress ONLY after successful write\n", + " current_index = i + 1\n", + "\n", + " # -------- GPU MEMORY CLEANUP --------\n", + " if torch.cuda.is_available():\n", + " torch.cuda.empty_cache()\n", + " torch.cuda.synchronize()\n", + "\n", + " # -------- SAVE LIKE YOUR DEEPSEEK VERSION --------\n", + " if (i + 1) % SAVE_INTERVAL == 0 or (i + 1) == len(df):\n", + " df.to_csv(output_file, index=False)\n", + " with open(index_file, \"w\") as f:\n", + " f.write(str(current_index))\n", + " print(f\"💾 Progress saved after row {i+1}\")\n", + "\n", + "# Final save\n", + "df.to_csv(output_file, index=False)\n", + "index_file.write_text(str(current_index))\n", + "print(\"✅ Finished full dataset.\")" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "a55a5e30-83f3-4f7c-a537-b1216d4e8a07", + "metadata": {}, + "outputs": [], + "source": [] + } + ], + "metadata": { + "kernelspec": { + "display_name": "pm-paper", + "language": "python", + "name": "pm-paper" + }, + "language_info": { + "codemirror_mode": { + "name": "ipython", + "version": 3 + }, + "file_extension": ".py", + "mimetype": "text/x-python", + "name": "python", + "nbconvert_exporter": "python", + "pygments_lexer": "ipython3", + "version": "3.11.13" + } + }, + "nbformat": 4, + "nbformat_minor": 5 +} diff --git a/jupyter_notebooks/.ipynb_checkpoints/Section_2-3-4_Figure_8_Step_2_response_comparison_and_consensus_extraction-checkpoint.ipynb b/jupyter_notebooks/.ipynb_checkpoints/Section_2-3-4_Figure_8_Step_2_response_comparison_and_consensus_extraction-checkpoint.ipynb new file mode 100644 index 0000000000000000000000000000000000000000..48219906a3bfb10d8098f05a77cb98aa389da34c --- /dev/null +++ b/jupyter_notebooks/.ipynb_checkpoints/Section_2-3-4_Figure_8_Step_2_response_comparison_and_consensus_extraction-checkpoint.ipynb @@ -0,0 +1,2862 @@ +{ + "cells": [ + { + "cell_type": "markdown", + "id": "b222ae9b-d94a-4892-9b7f-bf6a201ede47", + "metadata": {}, + "source": [ + "## Comprehensive Country Name Standardization with Fictional Place Detection\n", + "### Saves files as modelname_standardized_country.csv" + ] + }, + { + "cell_type": "code", + "execution_count": 12, + "id": "6cbaef9d-3058-4a59-a8ee-32fcc2062ed6", + "metadata": { + "execution": { + "iopub.execute_input": "2025-12-08T12:56:23.897161Z", + "iopub.status.busy": "2025-12-08T12:56:23.896726Z", + "iopub.status.idle": "2025-12-08T12:56:34.193068Z", + "shell.execute_reply": "2025-12-08T12:56:34.191636Z", + "shell.execute_reply.started": "2025-12-08T12:56:23.897125Z" + } + }, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "============================================================\n", + "COUNTRY STANDARDIZATION WITH SPECIAL MISTRAL HANDLING\n", + "============================================================\n", + "MISTRAL TRANSFORMATIONS:\n", + " - Remove brackets () from 'country' column\n", + " - Remove brackets () from 'full_name' column\n", + " - Convert 'actress' → 'actor' in profession\n", + " - Convert 'cosplayer' → 'online personality' in profession\n", + "\n", + "GEMMA & QWEN: No special transformations\n", + "============================================================\n", + "\n", + "Processing: gemma_local_annotated_POI.csv\n", + " - Found and standardized 'country' column\n", + " - Updated 'name' column for 4 fictional entries\n", + " - Updated 'real_name' column for 4 fictional entries\n", + " - Updated 'full_name' column for 4 fictional entries\n", + " - Changes: 6536 standardized, 4 fictional→Unknown\n", + " - Unique countries after standardization: ['Afghanistan', 'Albania', 'Algeria', 'American Samoa', 'Angola', 'Argentina', 'Armenia', 'Australia', 'Austria', 'Azerbaijan', 'Bahamas', 'Bangladesh', 'Barbados', 'Belarus', 'Belgium', 'Benin', 'Bolivia', 'Brazil', 'Bulgaria', 'Cambodia', 'Cameroon', 'Canada', 'Central African Republic', 'Chile', 'China', 'Colombia', 'Costa Rica', 'Croatia', 'Cuba', 'Cyprus', 'Czechia', 'Denmark', 'Dominican Republic', 'Egypt', 'El Salvador', 'Estonia', 'Ethiopia', 'Finland', 'France', 'French Polynesia', 'Georgia', 'Germany', 'Ghana', 'Greece', 'Guatemala', 'Guyana', 'Hong Kong', 'Hungary', 'Iceland', 'India', 'Indonesia', 'Iran', 'Iraq', 'Ireland', 'Israel', 'Italy', 'Jamaica', 'Japan', 'Jordan', 'Kazakhstan', 'Kenya', 'Kosovo', 'Kuwait', 'Kyrgyzstan', 'Latvia', 'Lebanon', 'Libya', 'Lithuania', 'Macau', 'Malaysia', 'Malta', 'Mexico', 'Moldova', 'Monaco', 'Mongolia', 'Morocco', 'Myanmar', 'Namibia', 'Nepal', 'Netherlands', 'New Zealand', 'Nicaragua', 'Nigeria', 'North Korea', 'North Macedonia', 'Norway', 'Pakistan', 'Palestine', 'Paraguay', 'Peru', 'Philippines', 'Poland', 'Portugal', 'Puerto Rico', 'Republic of the Congo', 'Romania', 'Russia', 'Samoa', 'Saudi Arabia', 'Senegal', 'Serbia', 'Singapore', 'Slovakia', 'Slovenia', 'Somalia', 'South Africa', 'South Korea', 'Spain', 'Sri Lanka', 'Sudan', 'Sweden', 'Switzerland', 'Syria', 'Taiwan', 'Tanzania', 'Thailand', 'Tonga', 'Trinidad and Tobago', 'Tunisia', 'Türkiye', 'UK', 'USA', 'Uganda', 'Ukraine', 'Unknown', 'Uruguay', 'Uzbekistan', 'Venezuela', 'Vietnam', 'Zambia', 'Zimbabwe']\n", + " - Saved to: gemma_standardized_country.csv\n", + " - Full path: /shares/weddigen.ki.uzh/laura_wagner/phase_01/pm-paper/data/CSV/gemma_standardized_country.csv\n", + "\n", + "Processing: mistral_local_annotated_POI.csv\n", + " ⚠️ MISTRAL DATA: Will remove bracketed content ()\n", + " - Updated 'name' column for fictional entries\n", + " - Saved to: mistral_standardized_country.csv\n", + " - Full path: /shares/weddigen.ki.uzh/laura_wagner/phase_01/pm-paper/data/CSV/mistral_standardized_country.csv\n", + "\n", + "Processing: qwen_local_annotated_POI.csv\n", + " - Found and standardized 'country' column\n", + " - Updated 'name' column for 14 fictional entries\n", + " - Updated 'real_name' column for 14 fictional entries\n", + " - Updated 'full_name' column for 14 fictional entries\n", + " - Changes: 26750 standardized, 14 fictional→Unknown\n", + " - Unique countries after standardization: ['Afghanistan', 'Albania', 'Algeria', 'American Samoa', 'Argentina', 'Armenia', 'Australia', 'Austria', 'Azerbaijan', 'Babylonian Empire', 'Bangladesh', 'Barbados', 'Belarus', 'Belgium', 'Benin', 'Bolivia', 'Bosnia and Herzegovina', 'Brazil', 'Bulgaria', 'Cameroon', 'Canada', 'Central African Republic', 'Chile', 'China', 'Colombia', 'Costa Rica', 'Croatia', 'Cuba', 'Cyprus', 'Czechia', 'Denmark', 'Discworld', 'Dominican Republic', 'Ecuador', 'Egypt', 'El Salvador', 'Estonia', 'Europe', 'Finland', 'France', 'French Polynesia', 'Georgia', 'Germany', 'Ghana', 'Greece', 'Guatemala', 'Haiti', 'Hong Kong', 'Hungary', 'Iceland', 'India', 'Indonesia', 'Iran', 'Iraq', 'Ireland', 'Israel', 'Italy', 'Jamaica', 'Japan', 'Jordan', 'Kazakhstan', 'Kenya', 'Kuwait', 'Kyrgyzstan', 'Laos', 'Latvia', 'Lebanon', 'Libya', 'Lithuania', 'Malaysia', 'Mali', 'Malta', 'Mexico', 'Moldova', 'Monaco', 'Mongolia', 'Morocco', 'Myanmar', 'Nepal', 'Netherlands', 'New Zealand', 'Nicaragua', 'Nigeria', 'North Korea', 'North Macedonia', 'Norway', 'Oman', 'Pakistan', 'Palestine', 'Papua New Guinea', 'Paraguay', 'Peru', 'Philippines', 'Poland', 'Portugal', 'Puerto Rico', 'Romania', 'Rome', 'Russia', 'Samoa', 'Saudi Arabia', 'Senegal', 'Serbia', 'Singapore', 'Slovakia', 'Slovenia', 'South Africa', 'South Korea', 'South Sudan', 'Spain', 'Sri Lanka', 'Sudan', 'Sweden', 'Switzerland', 'Syria', 'Taiwan', 'Thailand', 'Tunisia', 'Türkiye', 'UK', 'USA', 'Uganda', 'Ukraine', 'United Arab Emirates', 'Unknown', 'Uruguay', 'Uzbekistan', 'Vatican City', 'Venezuela', 'Vietnam', 'Worldwide', 'Zimbabwe']\n", + " - Saved to: qwen_standardized_country.csv\n", + " - Full path: /shares/weddigen.ki.uzh/laura_wagner/phase_01/pm-paper/data/CSV/qwen_standardized_country.csv\n", + "\n", + "============================================================\n", + "SUMMARY OF COUNTRY STANDARDIZATION\n", + "============================================================\n", + "\n", + "GEMMA:\n", + " Total rows: 50861\n", + " Top 10 countries:\n", + " - USA: 19378\n", + " - Unknown: 6320\n", + " - Japan: 4042\n", + " - UK: 2869\n", + " - South Korea: 2088\n", + " - China: 1765\n", + " - India: 1163\n", + " - Canada: 1150\n", + " - Russia: 1003\n", + " - France: 803\n", + "\n", + "MISTRAL:\n", + " Total rows: 50861\n", + " Top 10 countries:\n", + " - Unknown: 21033\n", + " - USA: 10564\n", + " - Japan: 3068\n", + " - UK: 2627\n", + " - South Korea: 2097\n", + " - China: 1409\n", + " - India: 1161\n", + " - Canada: 854\n", + " - France: 662\n", + " - Russia: 620\n", + "\n", + "QWEN:\n", + " Total rows: 50861\n", + " Top 10 countries:\n", + " - USA: 23017\n", + " - Japan: 4326\n", + " - UK: 3174\n", + " - Unknown: 2672\n", + " - China: 2403\n", + " - South Korea: 2333\n", + " - India: 1341\n", + " - Russia: 1036\n", + " - France: 851\n", + " - Canada: 828\n", + "\n", + "============================================================\n", + "Processing complete!\n", + "\n", + "Output files created:\n", + " - gemma_standardized_country.csv\n", + " - mistral_standardized_country.csv\n", + " - qwen_standardized_country.csv\n", + "\n", + "============================================================\n", + "MISTRAL-SPECIFIC TRANSFORMATIONS:\n", + " ✓ Country: Brackets () removed\n", + " ✓ Full_name: Brackets () removed\n", + " ✓ Profession: 'actress' → 'actor'\n", + " ✓ Profession: 'cosplayer' → 'online personality'\n", + "\n", + "GEMMA & QWEN:\n", + " ✓ No bracket removal\n", + " ✓ No profession transformations\n", + "============================================================\n" + ] + } + ], + "source": [ + "import pandas as pd\n", + "from pathlib import Path\n", + "import re\n", + "\n", + "# Define fictional places, companies, and invalid entries\n", + "FICTIONAL_PLACES = {\n", + " 'arrakis', 'asgard', 'krypton', 'gotham', 'metropolis', 'wakanda', \n", + " 'middle earth', 'hogwarts', 'westeros', 'narnia', 'mordor', 'gondor',\n", + " 'atlantis', 'shangri-la', 'eldorado', 'camelot', 'avalon', 'valhalla',\n", + " 'pandora', 'tatooine', 'coruscant', 'naboo', 'alderaan', 'endor',\n", + " 'rivendell', 'shire', 'rohan', 'isengard', 'minas tirith',\n", + " 'springfield', 'quahog', 'south park', 'bikini bottom',\n", + " 'emerald city', 'oz', 'neverland', 'wonderland'\n", + "}\n", + "\n", + "# Companies and brands that shouldn't be countries\n", + "COMPANIES_BRANDS = {\n", + " 'nintendo', 'hollywood', 'disney', 'pixar', 'marvel', 'dc',\n", + " 'sony', 'microsoft', 'apple', 'google', 'facebook', 'meta',\n", + " 'amazon', 'netflix', 'hbo', 'warner bros', 'paramount'\n", + "}\n", + "\n", + "# Historical entities that no longer exist\n", + "HISTORICAL_ENTITIES = {\n", + " 'roman republic', 'roman empire', 'byzantine empire', 'ottoman empire',\n", + " 'austro-hungarian empire', 'yugoslavia', 'east germany', 'west germany',\n", + " 'rhodesia', 'zaire', 'persia', 'prussia', 'mesopotamia', 'babylon'\n", + "}\n", + "\n", + "# Define country standardization mappings\n", + "COUNTRY_MAPPINGS = {\n", + " # USA variations\n", + " 'united states': 'USA',\n", + " 'united states of america': 'USA',\n", + " 'america': 'USA',\n", + " 'american': 'USA', # Added based on your data\n", + " 'us': 'USA',\n", + " 'u.s.': 'USA',\n", + " 'u.s.a.': 'USA',\n", + " 'states': 'USA',\n", + " 'unitedstates': 'USA',\n", + " \n", + " # UK variations\n", + " 'united kingdom': 'UK',\n", + " 'unitedkingdom': 'UK',\n", + " 'england': 'UK',\n", + " 'britain': 'UK',\n", + " 'great britain': 'UK',\n", + " 'uk': 'UK',\n", + " 'u.k.': 'UK',\n", + " 'scotland': 'UK', # Part of UK\n", + " 'wales': 'UK', # Part of UK\n", + " 'northern ireland': 'UK', # Part of UK\n", + " 'northernireland': 'UK',\n", + " \n", + " # Turkey -> Türkiye\n", + " 'turkey': 'Türkiye',\n", + " \n", + " # Czech Republic -> Czechia\n", + " 'czech republic': 'Czechia',\n", + " 'czechoslovakia': 'Czechia',\n", + " 'czechoslowakia': 'Czechia',\n", + " \n", + " # USSR -> Russia\n", + " 'ussr': 'Russia',\n", + " 'udssr': 'Russia',\n", + " 'soviet union': 'Russia',\n", + " \n", + " # Korea variations\n", + " 'korea': 'South Korea',\n", + " 'south korea': 'South Korea',\n", + " 'southkorea': 'South Korea',\n", + " 'republic of korea': 'South Korea',\n", + " 'north korea': 'North Korea',\n", + " 'northkorea': 'North Korea',\n", + " 'dprk': 'North Korea',\n", + " 'democratic people\\'s republic of korea': 'North Korea',\n", + " 'korea (democratic people\\'s republic of)': 'North Korea',\n", + " \n", + " # China variations\n", + " 'china': 'China',\n", + " 'people\\'s republic of china': 'China',\n", + " 'prc': 'China',\n", + " 'mainland china': 'China',\n", + " \n", + " # Hong Kong variations\n", + " 'hong kong': 'Hong Kong',\n", + " 'hongkong': 'Hong Kong',\n", + " \n", + " # Colombia/Columbia confusion\n", + " 'columbia': 'Colombia', # Common misspelling\n", + " 'colombia': 'Colombia',\n", + " \n", + " # Dominican Republic variations\n", + " 'dominican': 'Dominican Republic',\n", + " 'dominican republic': 'Dominican Republic',\n", + " 'dominicanrepublic': 'Dominican Republic',\n", + " \n", + " # Puerto Rico variations\n", + " 'puerto rico': 'Puerto Rico',\n", + " 'puertorico': 'Puerto Rico',\n", + " \n", + " # El Salvador variations\n", + " 'el salvador': 'El Salvador',\n", + " 'elsalvador': 'El Salvador',\n", + " \n", + " # South Africa variations\n", + " 'south africa': 'South Africa',\n", + " 'southafrica': 'South Africa',\n", + " \n", + " # Sri Lanka variations\n", + " 'sri lanka': 'Sri Lanka',\n", + " 'srilanka': 'Sri Lanka',\n", + " \n", + " # American Samoa variations\n", + " 'american samoa': 'American Samoa',\n", + " 'americansamoa': 'American Samoa',\n", + " \n", + " # North Macedonia variations\n", + " 'macedonia': 'North Macedonia',\n", + " 'northmacedonia': 'North Macedonia',\n", + " 'north macedonia': 'North Macedonia',\n", + " \n", + " # Central African Republic variations\n", + " 'central african republic': 'Central African Republic',\n", + " 'centralafricanrepublic': 'Central African Republic',\n", + " \n", + " # Republic of the Congo variations\n", + " 'republic of the congo': 'Republic of the Congo',\n", + " 'republic_of_the_congo': 'Republic of the Congo',\n", + " \n", + " # Common standardizations\n", + " 'holland': 'Netherlands',\n", + " 'the netherlands': 'Netherlands',\n", + " 'deutschland': 'Germany',\n", + " 'nippon': 'Japan',\n", + " 'espana': 'Spain',\n", + " 'españa': 'Spain',\n", + " \n", + " # Additional standardizations\n", + " 'vatican': 'Vatican City',\n", + " 'uae': 'United Arab Emirates',\n", + " 'emirates': 'United Arab Emirates',\n", + " 'bosnia': 'Bosnia and Herzegovina',\n", + " 'papua': 'Papua New Guinea',\n", + " 'png': 'Papua New Guinea',\n", + " 'trinidad': 'Trinidad and Tobago',\n", + "}\n", + "\n", + "def remove_bracketed_content(text):\n", + " \"\"\"\n", + " Remove everything in brackets (parentheses) from text.\n", + " Example: \"Japan (Asian country)\" -> \"Japan\"\n", + " \"\"\"\n", + " if pd.isna(text):\n", + " return text\n", + " \n", + " text_str = str(text).strip()\n", + " \n", + " # Remove content in parentheses including the parentheses\n", + " cleaned = re.sub(r'\\([^)]*\\)', '', text_str)\n", + " \n", + " # Strip any extra whitespace left after removal\n", + " return cleaned.strip()\n", + "\n", + "def standardize_profession(profession_value, is_mistral=False):\n", + " \"\"\"\n", + " Standardize profession values.\n", + " For Mistral data: actress -> actor, cosplayer -> online personality\n", + " \n", + " Args:\n", + " profession_value: The profession value to standardize\n", + " is_mistral: If True, applies Mistral-specific transformations\n", + " \"\"\"\n", + " if pd.isna(profession_value):\n", + " return profession_value\n", + " \n", + " profession_str = str(profession_value).strip()\n", + " \n", + " if not profession_str or profession_str == 'Unknown':\n", + " return profession_str\n", + " \n", + " # For Mistral data only, apply specific transformations\n", + " if is_mistral:\n", + " profession_lower = profession_str.lower()\n", + " \n", + " # actress -> actor\n", + " if profession_lower == 'actress':\n", + " return 'actor'\n", + " \n", + " # cosplayer -> online personality\n", + " if profession_lower == 'cosplayer':\n", + " return 'online personality'\n", + " \n", + " return profession_str\n", + "\n", + "def is_fictional_or_invalid(country_str):\n", + " \"\"\"\n", + " Check if a country string is fictional, a company, or historically invalid.\n", + " \"\"\"\n", + " country_lower = country_str.lower().strip()\n", + " \n", + " # Check for fictional places\n", + " if country_lower in FICTIONAL_PLACES:\n", + " return True\n", + " \n", + " # Check for companies/brands\n", + " if country_lower in COMPANIES_BRANDS:\n", + " return True\n", + " \n", + " # Check for historical entities\n", + " if country_lower in HISTORICAL_ENTITIES:\n", + " return True\n", + " \n", + " # Check for entries that contain \"Unknown\" with additional text\n", + " if 'unknown' in country_lower and len(country_lower) > 7:\n", + " return True\n", + " \n", + " # Check for entries that explicitly say \"fictional\"\n", + " if 'fictional' in country_lower:\n", + " return True\n", + " \n", + " # Check for entries with multiple countries (containing comma)\n", + " if ',' in country_str and country_lower != 'trinidad and tobago':\n", + " return True\n", + " \n", + " return False\n", + "\n", + "def standardize_country(country_value, is_mistral=False):\n", + " \"\"\"\n", + " Standardize a single country name based on the mapping.\n", + " \n", + " Args:\n", + " country_value: The country value to standardize\n", + " is_mistral: If True, removes bracketed content first (MISTRAL ONLY)\n", + " \"\"\"\n", + " if pd.isna(country_value):\n", + " return country_value\n", + " \n", + " # Convert to string and strip whitespace\n", + " country_str = str(country_value).strip()\n", + " \n", + " # For mistral data ONLY, remove bracketed content first\n", + " if is_mistral:\n", + " original_str = country_str\n", + " country_str = remove_bracketed_content(country_str)\n", + " # if original_str != country_str:\n", + " # print(f\" [Mistral] Removed brackets: '{original_str}' -> '{country_str}'\")\n", + " \n", + " # Return if empty or already 'Unknown'\n", + " if not country_str or country_str == 'Unknown':\n", + " return country_str if country_str else 'Unknown'\n", + " \n", + " # Check if it's fictional or invalid\n", + " if is_fictional_or_invalid(country_str):\n", + " return 'Unknown'\n", + " \n", + " # Convert to lowercase for matching\n", + " country_lower = country_str.lower()\n", + " \n", + " # Check if it matches any of our mappings\n", + " for pattern, replacement in COUNTRY_MAPPINGS.items():\n", + " if country_lower == pattern:\n", + " return replacement\n", + " \n", + " # If no exact match found, return original with proper capitalization\n", + " # This preserves valid countries not in our mapping\n", + " return country_str\n", + "\n", + "def extract_model_name(file_path):\n", + " \"\"\"\n", + " Extract the model name (gemma, mistral, qwen) from the file path.\n", + " \"\"\"\n", + " file_name = Path(file_path).stem.lower()\n", + " \n", + " if 'gemma' in file_name:\n", + " return 'gemma'\n", + " elif 'mistral' in file_name:\n", + " return 'mistral'\n", + " elif 'qwen' in file_name:\n", + " return 'qwen'\n", + " else:\n", + " # Fallback to using the full stem if model name not found\n", + " return Path(file_path).stem\n", + "\n", + "def process_csv_file(input_file, output_file):\n", + " \"\"\"\n", + " Process a CSV file to standardize country names and handle fictional places.\n", + " \"\"\"\n", + " # Convert to Path objects for consistent handling\n", + " input_path = Path(input_file)\n", + " output_path = Path(output_file)\n", + " \n", + " # Determine if this is mistral data\n", + " is_mistral = 'mistral' in input_path.name.lower()\n", + " \n", + " print(f\"Processing: {input_path.name}\")\n", + " if is_mistral:\n", + " print(f\" ⚠️ MISTRAL DATA: Will remove bracketed content ()\")\n", + " \n", + " # Track changes\n", + " changes_made = {'standardized': 0, 'fictional_to_unknown': 0, 'brackets_removed': 0}\n", + " \n", + " # For mistral.csv which might have no header, we need special handling\n", + " if is_mistral:\n", + " # Try to read normally first\n", + " try:\n", + " df = pd.read_csv(input_path)\n", + " # Check if 'country' column exists\n", + " if 'country' in df.columns:\n", + " original_values = df['country'].copy()\n", + " \n", + " # Count brackets before removal\n", + " changes_made['brackets_removed'] = original_values.astype(str).str.contains(r'\\(').sum()\n", + " \n", + " df['country'] = df['country'].apply(lambda x: standardize_country(x, is_mistral=True))\n", + " \n", + " # Count changes\n", + " changes_made['standardized'] = (original_values != df['country']).sum()\n", + " changes_made['fictional_to_unknown'] = ((df['country'] == 'Unknown') & (original_values != 'Unknown')).sum()\n", + " \n", + " #print(f\" - Found and standardized 'country' column\")\n", + " #print(f\" - Removed brackets from {changes_made['brackets_removed']} country entries\")\n", + " \n", + " # Also update 'name' column for fictional entries\n", + " if 'name' in df.columns:\n", + " fictional_mask = df['country'] == 'Unknown'\n", + " df.loc[fictional_mask & (original_values != 'Unknown'), 'name'] = 'Unknown'\n", + " print(f\" - Updated 'name' column for fictional entries\")\n", + " \n", + " # MISTRAL SPECIFIC: Remove brackets from full_name column\n", + " if 'full_name' in df.columns:\n", + " original_full_names = df['full_name'].copy()\n", + " brackets_in_names = original_full_names.astype(str).str.contains(r'\\(').sum()\n", + " df['full_name'] = df['full_name'].apply(remove_bracketed_content)\n", + " #print(f\" - Removed brackets from {brackets_in_names} full_name entries\")\n", + " \n", + " # MISTRAL SPECIFIC: Standardize profession_llm column\n", + " if 'profession_llm' in df.columns:\n", + " original_professions = df['profession_llm'].copy()\n", + " df['profession_llm'] = df['profession_llm'].apply(lambda x: standardize_profession(x, is_mistral=True))\n", + " \n", + " actress_count = (original_professions.astype(str).str.lower() == 'actress').sum()\n", + " cosplayer_count = (original_professions.astype(str).str.lower() == 'cosplayer').sum()\n", + " \n", + " if actress_count > 0:\n", + " print(f\" - Converted {actress_count} 'actress' → 'actor'\")\n", + " if cosplayer_count > 0:\n", + " print(f\" - Converted {cosplayer_count} 'cosplayer' → 'online personality'\")\n", + " \n", + " else:\n", + " # If no country column, assume last column\n", + " last_col = df.columns[-1]\n", + " df[last_col] = df[last_col].apply(lambda x: standardize_country(x, is_mistral=True))\n", + " print(f\" - Standardized column '{last_col}' (assumed to be country)\")\n", + " except:\n", + " # If normal reading fails, try without header\n", + " df = pd.read_csv(input_path, header=None)\n", + " last_col = df.columns[-1]\n", + " df[last_col] = df[last_col].apply(lambda x: standardize_country(x, is_mistral=True))\n", + " print(f\" - Standardized column {last_col} (assumed to be country, no header)\")\n", + " else:\n", + " # Normal CSV with header (GEMMA and QWEN - NO bracket removal)\n", + " df = pd.read_csv(input_path)\n", + " \n", + " # Check if 'country' column exists\n", + " if 'country' in df.columns:\n", + " original_values = df['country'].copy()\n", + " df['country'] = df['country'].apply(lambda x: standardize_country(x, is_mistral=False))\n", + " \n", + " # Count changes\n", + " changes_made['standardized'] = (original_values != df['country']).sum()\n", + " changes_made['fictional_to_unknown'] = ((df['country'] == 'Unknown') & (original_values != 'Unknown')).sum()\n", + " \n", + " print(f\" - Found and standardized 'country' column\")\n", + " \n", + " # Also update 'name' and 'real_name' columns for fictional entries\n", + " for col in ['name', 'real_name', 'full_name']:\n", + " if col in df.columns:\n", + " # Set to Unknown where country was changed to Unknown due to being fictional\n", + " fictional_mask = (df['country'] == 'Unknown') & (original_values != 'Unknown')\n", + " df.loc[fictional_mask, col] = 'Unknown'\n", + " if fictional_mask.any():\n", + " print(f\" - Updated '{col}' column for {fictional_mask.sum()} fictional entries\")\n", + " \n", + " print(f\" - Changes: {changes_made['standardized']} standardized, {changes_made['fictional_to_unknown']} fictional→Unknown\")\n", + " print(f\" - Unique countries after standardization: {sorted(df['country'].dropna().unique())}\")\n", + " else:\n", + " print(f\" - Warning: No 'country' column found in {input_path.name}\")\n", + " \n", + " # Save the processed file\n", + " df.to_csv(output_path, index=False)\n", + " print(f\" - Saved to: {output_path.name}\")\n", + " print(f\" - Full path: {output_path.absolute()}\\n\")\n", + " \n", + " return df\n", + "\n", + "# ============================================================\n", + "# MAIN EXECUTION\n", + "# ============================================================\n", + "\n", + "# Get current directory (assuming you're running from a notebook)\n", + "current_dir = Path.cwd()\n", + "\n", + "# Define your input files using Path objects\n", + "input_files = [\n", + " current_dir.parent / \"data/CSV/gemma_local_annotated_POI.csv\",\n", + " current_dir.parent / \"data/CSV/mistral_local_annotated_POI.csv\",\n", + " current_dir.parent / \"data/CSV/qwen_local_annotated_POI.csv\"\n", + "]\n", + "\n", + "# Create a results dictionary to store processed dataframes\n", + "results = {}\n", + "\n", + "print(\"=\"*60)\n", + "print(\"COUNTRY STANDARDIZATION WITH SPECIAL MISTRAL HANDLING\")\n", + "print(\"=\"*60)\n", + "print(\"MISTRAL TRANSFORMATIONS:\")\n", + "print(\" - Remove brackets () from 'country' column\")\n", + "print(\" - Remove brackets () from 'full_name' column\")\n", + "print(\" - Convert 'actress' → 'actor' in profession\")\n", + "print(\" - Convert 'cosplayer' → 'online personality' in profession\")\n", + "print(\"\\nGEMMA & QWEN: No special transformations\")\n", + "print(\"=\"*60 + \"\\n\")\n", + "\n", + "for input_path in input_files:\n", + " # Convert to Path object if it isn't already\n", + " input_path = Path(input_path)\n", + " \n", + " # Extract model name and create standardized output filename\n", + " model_name = extract_model_name(input_path)\n", + " output_path = input_path.parent / f\"{model_name}_standardized_country.csv\"\n", + " \n", + " # Check if input file exists\n", + " if input_path.exists():\n", + " df = process_csv_file(input_path, output_path)\n", + " # Use model name as key\n", + " results[model_name] = df\n", + " else:\n", + " print(f\"Warning: {input_path} not found! Please check the file path.\")\n", + " print(f\" Absolute path checked: {input_path.absolute()}\\n\")\n", + "\n", + "# Display summary statistics\n", + "print(\"=\" * 60)\n", + "print(\"SUMMARY OF COUNTRY STANDARDIZATION\")\n", + "print(\"=\" * 60)\n", + "\n", + "for name, df in results.items():\n", + " print(f\"\\n{name.upper()}:\")\n", + " print(f\" Total rows: {len(df)}\")\n", + " \n", + " # Find the country column\n", + " if 'mistral' in name.lower():\n", + " # For mistral, might need to check last column\n", + " if 'country' in df.columns:\n", + " country_col = 'country'\n", + " else:\n", + " country_col = df.columns[-1]\n", + " else:\n", + " country_col = 'country' if 'country' in df.columns else None\n", + " \n", + " if country_col is not None:\n", + " country_counts = df[country_col].value_counts()\n", + " print(f\" Top 10 countries:\")\n", + " for country, count in country_counts.head(10).items():\n", + " print(f\" - {country}: {count}\")\n", + "\n", + "print(\"\\n\" + \"=\" * 60)\n", + "print(\"Processing complete!\")\n", + "print(\"\\nOutput files created:\")\n", + "for model_name in results.keys():\n", + " print(f\" - {model_name}_standardized_country.csv\")\n", + "print(\"\\n\" + \"=\" * 60)\n", + "print(\"MISTRAL-SPECIFIC TRANSFORMATIONS:\")\n", + "print(\" ✓ Country: Brackets () removed\")\n", + "print(\" ✓ Full_name: Brackets () removed\")\n", + "print(\" ✓ Profession: 'actress' → 'actor'\")\n", + "print(\" ✓ Profession: 'cosplayer' → 'online personality'\")\n", + "print(\"\\nGEMMA & QWEN:\")\n", + "print(\" ✓ No bracket removal\")\n", + "print(\" ✓ No profession transformations\")\n", + "print(\"=\" * 60)" + ] + }, + { + "cell_type": "markdown", + "id": "b5e06b76-da90-4cff-b8dd-a2b27f56cc94", + "metadata": {}, + "source": [ + "## Create a combined CSV for visual comparison line by line" + ] + }, + { + "cell_type": "code", + "execution_count": 13, + "id": "7f2fe3fb-5863-45ec-9df3-9fa64751936b", + "metadata": { + "execution": { + "iopub.execute_input": "2025-12-08T12:56:43.814045Z", + "iopub.status.busy": "2025-12-08T12:56:43.813613Z", + "iopub.status.idle": "2025-12-08T12:56:50.203009Z", + "shell.execute_reply": "2025-12-08T12:56:50.201693Z", + "shell.execute_reply.started": "2025-12-08T12:56:43.814012Z" + } + }, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "Reading base file from gemma: gemma_standardized_country.csv\n", + " - Shape: (50861, 58)\n", + " - Columns: ['id', 'name', 'type', 'baseModel', 'downloadCount', 'nsfwLevel', 'modelVersions', 'publishedAt', 'usernameHash', 'downloadUrl']...\n", + "\n", + "Processing gemma: gemma_standardized_country.csv\n", + " - Added gemma_full_name\n", + " - Added gemma_aliases\n", + " - Added gemma_gender\n", + " - Added gemma_profession_llm\n", + " - Added gemma_country\n", + "\n", + "Processing mistral: mistral_standardized_country.csv\n", + " - Added mistral_full_name\n", + " - Added mistral_aliases\n", + " - Added mistral_gender\n", + " - Added mistral_profession_llm\n", + " - Added mistral_country\n", + "\n", + "Processing qwen: qwen_standardized_country.csv\n", + " - Added qwen_full_name\n", + " - Added qwen_aliases\n", + " - Added qwen_gender\n", + " - Added qwen_profession_llm\n", + " - Added qwen_country\n", + "\n", + "Reordering columns to group by field type...\n", + " - Added 3 columns for full_name\n", + " - Added 3 columns for aliases\n", + " - Added 3 columns for gender\n", + " - Added 3 columns for profession_llm\n", + " - Added 3 columns for country\n", + "\n", + "Combined CSV saved to: /shares/weddigen.ki.uzh/laura_wagner/phase_01/pm-paper/data/CSV/combined_llm_annotations.csv\n", + "Final shape: (50861, 67)\n", + "Column order: likely_country, gemma_full_name, mistral_full_name, qwen_full_name, gemma_aliases, mistral_aliases, qwen_aliases, gemma_gender, mistral_gender, qwen_gender, gemma_profession_llm, mistral_profession_llm, qwen_profession_llm, gemma_country, mistral_country...\n", + "\n", + "============================================================\n", + "SAMPLE COMPARISON ACROSS MODELS (first 5 rows)\n", + "============================================================\n", + "\n", + "FULL_NAME Comparison:\n", + "----------------------------------------\n", + " name gemma mistral qwen\n", + "0 IU Lee Ji-eun Lee Ji-eun Lee Ji-eun\n", + "1 Super Pose Book Vol.1 - ControlNet Unknown Unknown Super Pose Book\n", + "2 Liyuu LoRA Li Yuchan Liyuu Li Yuxiao\n", + "3 Irene Irene Kim Irene Irene Kim\n", + "4 AESPA Karina Yoo Jimin Karina Kim Karina\n", + "\n", + "COUNTRY Comparison:\n", + "----------------------------------------\n", + " name gemma mistral qwen\n", + "0 IU South Korea South Korea South Korea\n", + "1 Super Pose Book Vol.1 - ControlNet Unknown Unknown Unknown\n", + "2 Liyuu LoRA China Unknown China\n", + "3 Irene South Korea Unknown South Korea\n", + "4 AESPA Karina South Korea South Korea South Korea\n", + "\n", + "PROFESSION_LLM Comparison:\n", + "----------------------------------------\n", + " name gemma mistral qwen\n", + "0 IU singer/musician, actor, tv personality singer/musician, tv personality, actress singer/musician, model, public figure\n", + "1 Super Pose Book Vol.1 - ControlNet model, adult performer, online personality model, online personality, actress model, online personality, artist\n", + "2 Liyuu LoRA singer/musician, actor, model model, online personality, artist model, online personality\n", + "3 Irene model, online personality, actor actor, model, tv personality singer/musician, model, public figure\n", + "4 AESPA Karina singer/musician, tv personality, online personality singer/musician, tv personality, public figure singer/musician, model, online personality\n", + "\n", + "============================================================\n", + "SUMMARY STATISTICS\n", + "============================================================\n", + "\n", + "GEMMA Country Distribution:\n", + "gemma_country\n", + "USA 19378\n", + "Unknown 6320\n", + "Japan 4042\n", + "UK 2869\n", + "South Korea 2088\n", + "Name: count, dtype: int64\n", + "\n", + "MISTRAL Country Distribution:\n", + "mistral_country\n", + "Unknown 21033\n", + "USA 10564\n", + "Japan 3068\n", + "UK 2627\n", + "South Korea 2097\n", + "Name: count, dtype: int64\n", + "\n", + "QWEN Country Distribution:\n", + "qwen_country\n", + "USA 23017\n", + "Japan 4326\n", + "UK 3174\n", + "Unknown 2672\n", + "China 2403\n", + "Name: count, dtype: int64\n", + "\n", + "============================================================\n", + "FINAL COLUMNS IN COMBINED FILE (GROUPED BY TYPE):\n", + "============================================================\n", + "\n", + "Common columns:\n", + " - id\n", + " - name\n", + " - type\n", + " - baseModel\n", + " - downloadCount\n", + " - nsfwLevel\n", + " - modelVersions\n", + " - publishedAt\n", + " - usernameHash\n", + " - downloadUrl\n", + " - firstImageUrl\n", + " - latestImageUrl\n", + " - poi\n", + " - AutoV1\n", + " - AutoV2\n", + " - AutoV3\n", + " - SHA256\n", + " - CRC32\n", + " - BLAKE3\n", + " - previewImage\n", + " ... and 32 more\n", + "\n", + "Full name columns:\n", + " - gemma_full_name\n", + " - mistral_full_name\n", + " - qwen_full_name\n", + "\n", + "Aliases columns:\n", + " - gemma_aliases\n", + " - mistral_aliases\n", + " - qwen_aliases\n", + "\n", + "Gender columns:\n", + " - gemma_gender\n", + " - mistral_gender\n", + " - qwen_gender\n", + "\n", + "Profession llm columns:\n", + " - gemma_profession_llm\n", + " - mistral_profession_llm\n", + " - qwen_profession_llm\n", + "\n", + "Country columns:\n", + " - likely_country\n", + " - gemma_country\n", + " - mistral_country\n", + " - qwen_country\n", + "\n", + "============================================================\n", + "✓ Combined file saved to: combined_llm_annotations.csv\n", + " Total rows: 50861\n", + " Total columns: 67\n", + "============================================================\n" + ] + } + ], + "source": [ + "# Script to combine LLM-annotated CSVs with model-specific columns\n", + "# Creates a single CSV with columns grouped by field type\n", + "\n", + "import pandas as pd\n", + "from pathlib import Path\n", + "\n", + "def combine_llm_annotations(input_files, output_file):\n", + " \"\"\"\n", + " Combine multiple LLM-annotated CSVs into one file with model-specific columns.\n", + " Columns are grouped by field type (full_name, aliases, etc.) across models.\n", + " \n", + " Args:\n", + " input_files: Dictionary with model names as keys and file paths as values\n", + " output_file: Path to save the combined CSV\n", + " \"\"\"\n", + " \n", + " # LLM-specific columns that will be prefixed with model name\n", + " llm_columns = ['full_name', 'aliases', 'gender', 'profession_llm', 'country']\n", + " \n", + " # Columns that should be kept from the first file (common across all)\n", + " # Adjust this list based on your actual columns\n", + " common_columns = [\n", + " 'id', 'name', 'type', 'baseModel', 'downloadCount', 'nsfwLevel', \n", + " 'modelVersions', 'publishedAt', 'usernameHash', 'downloadUrl',\n", + " 'firstImageUrl', 'latestImageUrl', 'poi', 'AutoV1', 'AutoV2', 'AutoV3',\n", + " 'SHA256', 'CRC32', 'BLAKE3', 'previewImage',\n", + " 'version_id_1', 'version_id_2', 'version_id_3', 'version_id_4', 'version_id_5',\n", + " 'version_id_6', 'version_id_7', 'version_id_8', 'version_id_9', 'version_id_10',\n", + " 'version_id_11', 'version_id_12', 'version_id_13', 'version_id_14', 'version_id_15',\n", + " 'version_id_16', 'version_id_17', 'version_id_18', 'version_id_19', 'version_id_20',\n", + " 'tag_1', 'tag_2', 'tag_3', 'tag_4', 'tag_5', 'tag_6', 'tag_7',\n", + " 'real_name', 'tags', 'likely_country', 'likely_nationality', 'likely_profession'\n", + " ]\n", + " \n", + " # Read the first file to get the base dataframe\n", + " first_model = list(input_files.keys())[0]\n", + " first_path = input_files[first_model]\n", + " \n", + " print(f\"Reading base file from {first_model}: {first_path.name}\")\n", + " \n", + " # Read the first file\n", + " if 'mistral' in first_model and not pd.read_csv(first_path, nrows=1).columns.str.contains('id').any():\n", + " # Special handling for mistral if it has no header\n", + " base_df = pd.read_csv(first_path, header=None)\n", + " # We'll need to map columns manually\n", + " print(\" - Note: Mistral file appears to have no header, handling specially\")\n", + " else:\n", + " base_df = pd.read_csv(first_path)\n", + " \n", + " print(f\" - Shape: {base_df.shape}\")\n", + " print(f\" - Columns: {base_df.columns.tolist()[:10]}...\")\n", + " \n", + " # Start with common columns that exist in the base dataframe\n", + " existing_common_cols = [col for col in common_columns if col in base_df.columns]\n", + " result_df = base_df[existing_common_cols].copy()\n", + " \n", + " # Dictionary to collect columns by type\n", + " columns_by_type = {col_type: [] for col_type in llm_columns}\n", + " \n", + " # Process each model's data\n", + " for model_name, file_path in input_files.items():\n", + " print(f\"\\nProcessing {model_name}: {file_path.name}\")\n", + " \n", + " # Read the model's data\n", + " if 'mistral' in model_name and not pd.read_csv(file_path, nrows=1).columns.str.contains('id').any():\n", + " df = pd.read_csv(file_path, header=None)\n", + " # Map columns based on position (you may need to adjust this)\n", + " # Assuming the last 5 columns are the LLM annotations\n", + " num_cols = len(df.columns)\n", + " llm_start_idx = num_cols - 5\n", + " \n", + " for i, col_name in enumerate(llm_columns):\n", + " if llm_start_idx + i < num_cols:\n", + " new_col_name = f'{model_name}_{col_name}'\n", + " result_df[new_col_name] = df.iloc[:, llm_start_idx + i]\n", + " columns_by_type[col_name].append(new_col_name)\n", + " print(f\" - Added {new_col_name} from column {llm_start_idx + i}\")\n", + " else:\n", + " df = pd.read_csv(file_path)\n", + " \n", + " # Add LLM-specific columns with model prefix\n", + " for col in llm_columns:\n", + " if col in df.columns:\n", + " new_col_name = f'{model_name}_{col}'\n", + " result_df[new_col_name] = df[col]\n", + " columns_by_type[col].append(new_col_name)\n", + " print(f\" - Added {new_col_name}\")\n", + " else:\n", + " print(f\" - Warning: Column '{col}' not found in {model_name}\")\n", + " \n", + " # Reorder columns to group by type\n", + " print(\"\\nReordering columns to group by field type...\")\n", + " \n", + " # Get the final column order\n", + " final_columns = existing_common_cols.copy()\n", + " \n", + " # Add columns grouped by type\n", + " for col_type in llm_columns:\n", + " if columns_by_type[col_type]:\n", + " final_columns.extend(columns_by_type[col_type])\n", + " print(f\" - Added {len(columns_by_type[col_type])} columns for {col_type}\")\n", + " \n", + " # Reorder the dataframe\n", + " result_df = result_df[final_columns]\n", + " \n", + " # Save the combined dataframe\n", + " result_df.to_csv(output_file, index=False)\n", + " print(f\"\\nCombined CSV saved to: {output_file}\")\n", + " print(f\"Final shape: {result_df.shape}\")\n", + " print(f\"Column order: {', '.join([col for col in result_df.columns if any(col.endswith(f'_{ctype}') for ctype in llm_columns)][:15])}...\")\n", + " \n", + " return result_df\n", + "\n", + "def get_column_comparison(df, models):\n", + " \"\"\"\n", + " Create a comparison of values across models for inspection.\n", + " \"\"\"\n", + " print(\"\\n\" + \"=\"*60)\n", + " print(\"SAMPLE COMPARISON ACROSS MODELS (first 5 rows)\")\n", + " print(\"=\"*60)\n", + " \n", + " # Show a few examples comparing across models\n", + " comparison_cols = ['full_name', 'country', 'profession_llm']\n", + " \n", + " for col in comparison_cols:\n", + " print(f\"\\n{col.upper()} Comparison:\")\n", + " print(\"-\" * 40)\n", + " \n", + " # Create a comparison dataframe\n", + " comp_data = {'name': df['name'].iloc[:5] if 'name' in df.columns else df.index[:5]}\n", + " \n", + " for model in models:\n", + " col_name = f'{model}_{col}'\n", + " if col_name in df.columns:\n", + " comp_data[model] = df[col_name].iloc[:5]\n", + " \n", + " comp_df = pd.DataFrame(comp_data)\n", + " print(comp_df.to_string())\n", + "\n", + "# ============================================================\n", + "# MAIN EXECUTION\n", + "# ============================================================\n", + "\n", + "# Get current directory\n", + "current_dir = Path.cwd()\n", + "\n", + "# Define input files - these should be your standardized files\n", + "# Adjust the paths as needed\n", + "input_files = {\n", + " 'gemma': current_dir.parent / \"data/CSV/gemma_standardized_country.csv\",\n", + " 'mistral': current_dir.parent / \"data/CSV/mistral_standardized_country.csv\",\n", + " 'qwen': current_dir.parent / \"data/CSV/qwen_standardized_country.csv\"\n", + "}\n", + "\n", + "# Check if all files exist\n", + "all_exist = True\n", + "for model, path in input_files.items():\n", + " if not path.exists():\n", + " print(f\"Warning: {path} does not exist!\")\n", + " all_exist = False\n", + "\n", + "if all_exist:\n", + " # Define output file\n", + " output_file = current_dir.parent / \"data/CSV/combined_llm_annotations.csv\"\n", + " \n", + " # Combine the files\n", + " combined_df = combine_llm_annotations(input_files, output_file)\n", + " \n", + " # Show comparison\n", + " get_column_comparison(combined_df, list(input_files.keys()))\n", + " \n", + " # Show summary statistics\n", + " print(\"\\n\" + \"=\"*60)\n", + " print(\"SUMMARY STATISTICS\")\n", + " print(\"=\"*60)\n", + " \n", + " for model in input_files.keys():\n", + " country_col = f'{model}_country'\n", + " if country_col in combined_df.columns:\n", + " print(f\"\\n{model.upper()} Country Distribution:\")\n", + " print(combined_df[country_col].value_counts().head(5))\n", + " \n", + " print(\"\\n\" + \"=\"*60)\n", + " print(\"FINAL COLUMNS IN COMBINED FILE (GROUPED BY TYPE):\")\n", + " print(\"=\"*60)\n", + " \n", + " # Group columns by type for better display\n", + " common_cols = [col for col in combined_df.columns if not any(col.startswith(m + '_') for m in input_files.keys())]\n", + " \n", + " print(\"\\nCommon columns:\")\n", + " for col in common_cols[:20]: # Show first 20 common columns\n", + " print(f\" - {col}\")\n", + " if len(common_cols) > 20:\n", + " print(f\" ... and {len(common_cols)-20} more\")\n", + " \n", + " # Show columns grouped by type\n", + " llm_columns = ['full_name', 'aliases', 'gender', 'profession_llm', 'country']\n", + " for col_type in llm_columns:\n", + " type_cols = [col for col in combined_df.columns if col.endswith(f'_{col_type}')]\n", + " if type_cols:\n", + " print(f\"\\n{col_type.capitalize().replace('_', ' ')} columns:\")\n", + " for col in type_cols:\n", + " print(f\" - {col}\")\n", + " \n", + " print(\"\\n\" + \"=\"*60)\n", + " print(f\"✓ Combined file saved to: {output_file.name}\")\n", + " print(f\" Total rows: {len(combined_df)}\")\n", + " print(f\" Total columns: {len(combined_df.columns)}\")\n", + " print(\"=\"*60)\n", + "\n", + "else:\n", + " print(\"\\nPlease ensure all standardized files exist before running this script.\")\n", + " print(\"Expected files:\")\n", + " for model, path in input_files.items():\n", + " print(f\" - {path}\")" + ] + }, + { + "cell_type": "markdown", + "id": "ae787ec8-7403-492d-8552-026f95d9e5b9", + "metadata": {}, + "source": [ + "### Strict version" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "ca58ae0f-a3e1-44c9-9645-412997f1777d", + "metadata": { + "execution": { + "iopub.execute_input": "2025-12-08T13:00:58.348391Z", + "iopub.status.busy": "2025-12-08T13:00:58.347965Z", + "iopub.status.idle": "2025-12-08T13:01:02.440993Z", + "shell.execute_reply": "2025-12-08T13:01:02.439833Z", + "shell.execute_reply.started": "2025-12-08T13:00:58.348357Z" + } + }, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "================================================================================\n", + "IMPROVED CONSENSUS CREATION\n", + "================================================================================\n", + "Method: hybrid\n", + "Reading: /home/lauhp/000_PHD/000_010_PUBLICATION/CODE/pm-paper/data/CSV/combined_llm_annotations.csv\n", + "Input shape: (50861, 67)\n", + "Models: gemma, mistral, qwen\n", + "\n", + "Processing rows...\n", + " Processed 50000/50861 rows...\n", + " Processed 50861 rows. \n", + "\n", + "================================================================================\n", + "RESULTS\n", + "================================================================================\n", + "Total input rows: 50,861\n", + "Rows passing all criteria: 23,084 (45.4%)\n", + "\n", + "Consensus method usage:\n", + " - weighted: 43,445 (188.2%)\n", + " - adult_special: 5,622 (24.4%)\n", + " - adult_any_position: 1,784 (7.7%)\n", + " - no_consensus: 10 (0.0%)\n", + "\n", + "================================================================================\n", + "PROFESSION DISTRIBUTION\n", + "================================================================================\n", + "\n", + "Top 20 professions:\n", + "consensus_profession\n", + "actor 11046\n", + "singer/musician 3939\n", + "model 3017\n", + "online personality 1651\n", + "adult performer 1206\n", + "public figure 986\n", + "sports professional 458\n", + "voice actor/asmr 376\n", + "tv personality 373\n", + "wrestler 10\n", + "comedian 8\n", + "cheerleader 2\n", + "actress 2\n", + "dancer 2\n", + "architect 1\n", + "entrepreneur 1\n", + "basketball player 1\n", + "podcaster 1\n", + "artist 1\n", + "gymnast 1\n", + "Name: count, dtype: int64\n", + "\n", + "🎯 Adult performer variants: 1206 (5.22%)\n", + "\n", + "================================================================================\n", + "✓ Improved consensus saved to: improved_consensus.csv\n", + " Total rows: 23,084\n", + "================================================================================\n", + "\n", + "✅ Complete!\n" + ] + } + ], + "source": [ + "# Script to create strict consensus CSV\n", + "# Only keeps rows where:\n", + "# - All 3 models agree on country\n", + "# - All 3 models agree on gender\n", + "# - At least 2 models agree on profession\n", + "\n", + "import pandas as pd\n", + "from pathlib import Path\n", + "from collections import Counter\n", + "\n", + "def normalize_value(value):\n", + " \"\"\"Normalize values for comparison (handle NaN, whitespace, case)\"\"\"\n", + " if pd.isna(value):\n", + " return None\n", + " return str(value).strip().lower()\n", + "\n", + "def is_unknown_value(value):\n", + " \"\"\"Check if a value represents 'unknown' or similar non-informative values\"\"\"\n", + " if value is None:\n", + " return True\n", + " \n", + " value_str = str(value).strip().lower()\n", + " \n", + " # List of patterns that indicate unknown/missing data\n", + " unknown_patterns = [\n", + " 'unknown',\n", + " 'n/a',\n", + " 'na',\n", + " 'none',\n", + " 'not specified',\n", + " 'not available',\n", + " 'unclear',\n", + " 'uncertain',\n", + " '',\n", + " 'null'\n", + " ]\n", + " \n", + " return value_str in unknown_patterns\n", + "\n", + "def get_consensus_value(values, required_agreement=2):\n", + " \"\"\"\n", + " Get consensus value if at least 'required_agreement' models agree.\n", + " Returns (consensus_value, count_agreeing, all_values_dict)\n", + " \"\"\"\n", + " # Normalize values\n", + " normalized = [normalize_value(v) for v in values]\n", + " \n", + " # Count occurrences (excluding None)\n", + " valid_values = [v for v in normalized if v is not None]\n", + " \n", + " if not valid_values:\n", + " return None, 0, {}\n", + " \n", + " value_counts = Counter(valid_values)\n", + " most_common_value, count = value_counts.most_common(1)[0]\n", + " \n", + " # Create a dict of all values for inspection\n", + " all_values = {f\"model_{i+1}\": values[i] for i in range(len(values))}\n", + " \n", + " if count >= required_agreement:\n", + " return most_common_value, count, all_values\n", + " else:\n", + " return None, count, all_values\n", + "\n", + "def create_strict_consensus(input_file, output_file, models=['gemma', 'mistral', 'qwen']):\n", + " \"\"\"\n", + " Create a strict consensus CSV with only rows where:\n", + " - All 3 models agree on country\n", + " - All 3 models agree on gender\n", + " - At least 2 models agree on profession\n", + " \"\"\"\n", + " print(\"=\"*60)\n", + " print(\"STRICT CONSENSUS FILTER\")\n", + " print(\"=\"*60)\n", + " print(f\"Reading: {input_file}\")\n", + " \n", + " # Read the combined file\n", + " df = pd.read_csv(input_file)\n", + " \n", + " print(f\"Input shape: {df.shape}\")\n", + " print(f\"Models to check: {', '.join(models)}\")\n", + " \n", + " # Initialize lists to store results\n", + " rows_to_keep = []\n", + " consensus_data = []\n", + " \n", + " # Stats tracking\n", + " stats = {\n", + " 'total_rows': len(df),\n", + " 'country_fail': 0,\n", + " 'gender_fail': 0,\n", + " 'profession_fail': 0,\n", + " 'unknown_values': 0,\n", + " 'all_pass': 0\n", + " }\n", + " \n", + " print(\"\\nProcessing rows...\")\n", + " \n", + " for idx, row in df.iterrows():\n", + " if idx % 1000 == 0:\n", + " print(f\" Processed {idx}/{len(df)} rows...\", end='\\r')\n", + " \n", + " # Get values for each field from all models\n", + " countries = [row[f'{model}_country'] for model in models]\n", + " genders = [row[f'{model}_gender'] for model in models]\n", + " professions = [row[f'{model}_profession_llm'] for model in models]\n", + " \n", + " # Check country consensus (all 3 must agree)\n", + " country_consensus, country_count, country_vals = get_consensus_value(countries, required_agreement=3)\n", + " \n", + " # Check gender consensus (all 3 must agree)\n", + " gender_consensus, gender_count, gender_vals = get_consensus_value(genders, required_agreement=3)\n", + " \n", + " # Check profession consensus (at least 2 must agree)\n", + " profession_consensus, profession_count, profession_vals = get_consensus_value(professions, required_agreement=2)\n", + " \n", + " # Determine if row passes all criteria\n", + " country_pass = country_count == 3\n", + " gender_pass = gender_count == 3\n", + " profession_pass = profession_count >= 2\n", + " \n", + " # Check if any consensus value is \"unknown\" or similar\n", + " has_unknown = (\n", + " is_unknown_value(country_consensus) or \n", + " is_unknown_value(gender_consensus) or \n", + " is_unknown_value(profession_consensus)\n", + " )\n", + " \n", + " # Track failures\n", + " if not country_pass:\n", + " stats['country_fail'] += 1\n", + " if not gender_pass:\n", + " stats['gender_fail'] += 1\n", + " if not profession_pass:\n", + " stats['profession_fail'] += 1\n", + " if has_unknown:\n", + " stats['unknown_values'] += 1\n", + " \n", + " # Only keep row if all criteria pass AND no unknown values\n", + " if country_pass and gender_pass and profession_pass and not has_unknown:\n", + " rows_to_keep.append(idx)\n", + " stats['all_pass'] += 1\n", + " \n", + " consensus_data.append({\n", + " 'consensus_country': country_consensus,\n", + " 'consensus_gender': gender_consensus,\n", + " 'consensus_profession': profession_consensus,\n", + " 'profession_agreement_count': profession_count\n", + " })\n", + " \n", + " print(f\"\\n Processed {len(df)} rows. \")\n", + " \n", + " # Create consensus dataframe\n", + " if rows_to_keep:\n", + " # Get the original data for kept rows\n", + " result_df = df.iloc[rows_to_keep].copy().reset_index(drop=True)\n", + " \n", + " # Add consensus columns at the beginning (after common columns)\n", + " consensus_df = pd.DataFrame(consensus_data)\n", + " \n", + " # Find where to insert consensus columns (after 'name' if it exists, otherwise at start)\n", + " if 'name' in result_df.columns:\n", + " name_idx = result_df.columns.get_loc('name') + 1\n", + " else:\n", + " name_idx = 0\n", + " \n", + " # Insert consensus columns\n", + " for i, col in enumerate(consensus_df.columns):\n", + " result_df.insert(name_idx + i, col, consensus_df[col])\n", + " \n", + " # Save the result\n", + " result_df.to_csv(output_file, index=False)\n", + " \n", + " print(\"\\n\" + \"=\"*60)\n", + " print(\"FILTERING RESULTS\")\n", + " print(\"=\"*60)\n", + " print(f\"Total input rows: {stats['total_rows']:,}\")\n", + " print(f\"Rows passing all criteria: {stats['all_pass']:,} ({stats['all_pass']/stats['total_rows']*100:.1f}%)\")\n", + " print(f\"\\nFailure reasons (rows can fail multiple):\")\n", + " print(f\" - Country disagreement: {stats['country_fail']:,} ({stats['country_fail']/stats['total_rows']*100:.1f}%)\")\n", + " print(f\" - Gender disagreement: {stats['gender_fail']:,} ({stats['gender_fail']/stats['total_rows']*100:.1f}%)\")\n", + " print(f\" - Profession disagreement: {stats['profession_fail']:,} ({stats['profession_fail']/stats['total_rows']*100:.1f}%)\")\n", + " print(f\" - Contains 'Unknown': {stats['unknown_values']:,} ({stats['unknown_values']/stats['total_rows']*100:.1f}%)\")\n", + " \n", + " print(\"\\n\" + \"=\"*60)\n", + " print(\"CONSENSUS DISTRIBUTIONS\")\n", + " print(\"=\"*60)\n", + " \n", + " print(\"\\nCountry (top 10):\")\n", + " print(result_df['consensus_country'].value_counts().head(10))\n", + " \n", + " print(\"\\nGender:\")\n", + " print(result_df['consensus_gender'].value_counts())\n", + " \n", + " print(\"\\nProfession (top 10):\")\n", + " print(result_df['consensus_profession'].value_counts().head(10))\n", + " \n", + " print(\"\\nProfession agreement level:\")\n", + " print(result_df['profession_agreement_count'].value_counts().sort_index())\n", + " \n", + " print(\"\\n\" + \"=\"*60)\n", + " print(f\"✓ Strict consensus file saved to: {output_file.name}\")\n", + " print(f\" Total rows: {len(result_df):,}\")\n", + " print(f\" Total columns: {len(result_df.columns)}\")\n", + " print(\"=\"*60)\n", + " \n", + " return result_df\n", + " else:\n", + " print(\"\\n⚠ WARNING: No rows passed all criteria!\")\n", + " print(\"Creating empty file with proper columns...\")\n", + " \n", + " # Create empty dataframe with proper structure\n", + " result_df = df.iloc[:0].copy()\n", + " consensus_cols = ['consensus_country', 'consensus_gender', 'consensus_profession', 'profession_agreement_count']\n", + " for col in consensus_cols:\n", + " result_df.insert(0, col, [])\n", + " \n", + " result_df.to_csv(output_file, index=False)\n", + " return result_df\n", + "\n", + "def show_sample_comparisons(df, models, n_samples=5):\n", + " \"\"\"Show sample comparisons between models for quality check\"\"\"\n", + " print(\"\\n\" + \"=\"*60)\n", + " print(f\"SAMPLE COMPARISONS (first {n_samples} rows)\")\n", + " print(\"=\"*60)\n", + " \n", + " if len(df) == 0:\n", + " print(\"No data to display\")\n", + " return\n", + " \n", + " sample_df = df.head(n_samples)\n", + " \n", + " for idx, row in sample_df.iterrows():\n", + " print(f\"\\n--- Row {idx + 1}: {row.get('name', 'N/A')} ---\")\n", + " \n", + " # Country comparison\n", + " print(\"Country:\")\n", + " print(f\" Consensus: {row['consensus_country']}\")\n", + " for model in models:\n", + " col = f'{model}_country'\n", + " if col in row:\n", + " print(f\" {model}: {row[col]}\")\n", + " \n", + " # Gender comparison\n", + " print(\"Gender:\")\n", + " print(f\" Consensus: {row['consensus_gender']}\")\n", + " for model in models:\n", + " col = f'{model}_gender'\n", + " if col in row:\n", + " print(f\" {model}: {row[col]}\")\n", + " \n", + " # Profession comparison\n", + " print(f\"Profession (agreement: {row['profession_agreement_count']}/3):\")\n", + " print(f\" Consensus: {row['consensus_profession']}\")\n", + " for model in models:\n", + " col = f'{model}_profession_llm'\n", + " if col in row:\n", + " print(f\" {model}: {row[col]}\")\n", + "\n", + "# ============================================================\n", + "# MAIN EXECUTION\n", + "# ============================================================\n", + "\n", + "if __name__ == \"__main__\":\n", + " # Get current directory\n", + " current_dir = Path.cwd()\n", + " \n", + " # Define input file (the combined file from previous script)\n", + " input_file = current_dir.parent / \"data/CSV/combined_llm_annotations.csv\"\n", + " \n", + " # Define output file\n", + " output_file = current_dir.parent / \"data/CSV/strict_consensus.csv\"\n", + " \n", + " # Check if input file exists\n", + " if not input_file.exists():\n", + " print(f\"Error: Input file not found: {input_file}\")\n", + " print(\"Please run the combine script first to create combined_llm_annotations.csv\")\n", + " else:\n", + " # Models to check\n", + " models = ['gemma', 'mistral', 'qwen']\n", + " \n", + " # Create strict consensus file\n", + " result_df = create_strict_consensus(input_file, output_file, models)\n", + " \n", + " # Show sample comparisons\n", + " if len(result_df) > 0:\n", + " show_sample_comparisons(result_df, models, n_samples=5)\n", + " \n", + " print(\"\\n\" + \"=\"*60)\n", + " print(\"COMPLETE!\")\n", + " print(\"=\"*60)" + ] + }, + { + "cell_type": "markdown", + "id": "e0f8f331-a6f3-49e8-9734-ee7ffafebabc", + "metadata": {}, + "source": [ + "# ANALYSIS OF MODEL AGREEMENT WITH REFINED NAME MATCHING\n", + "\n", + "RULES FOR VALIDATION:\n", + "1. A row is considered \"valid\" if at least two models agree on:\n", + " - Country: Exact match required\n", + " - Gender: Exact match required \n", + " - Profession: Values are comma-separated lists; considered matching if:\n", + " a) They have at least two professions in common (order doesn'tmatter)\n", + "\n", + "2. For names (full_name) - REFINED RULES:\n", + " - Case-insensitive comparison\n", + " - Remove content in parentheses/brackets before comparison\n", + " - Normalize special characters and spacing\n", + " - Split into tokens and apply flexible matching:\n", + " * Minimum threshold: At least 50% of tokens must match\n", + " * Allow variations like \"Olivia Rose Holt\", \"Olivia Hastings Holt\", \"Olivia Holt\"\n", + " * Handle name variations like \"Kim Na-eun\" vs \"Kim Naeun\"\n", + " * Consider partial matches when most tokens align\n", + "\n", + "3. All Three Models Agreement:\n", + " - All three models must agree on country, gender, and have matching professions\n", + " - For names: All three must meet the refined matching criteria with each other" + ] + }, + { + "cell_type": "code", + "execution_count": 4, + "id": "fe38cf23-771d-4888-a5ed-f2b8beb17017", + "metadata": { + "execution": { + "iopub.execute_input": "2025-12-08T10:58:11.949915Z", + "iopub.status.busy": "2025-12-08T10:58:11.949461Z", + "iopub.status.idle": "2025-12-08T10:58:29.470022Z", + "shell.execute_reply": "2025-12-08T10:58:29.468836Z", + "shell.execute_reply.started": "2025-12-08T10:58:11.949880Z" + } + }, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "Loading combined data from: /shares/weddigen.ki.uzh/laura_wagner/phase_01/pm-paper/data/CSV/combined_llm_annotations.csv\n", + "Loaded 50,861 rows with 67 columns\n", + "Analyzing agreement between models: gemma, qwen, mistral\n", + "Total rows to analyze: 50861\n", + "\n", + "============================================================\n", + "AGREEMENT ANALYSIS RESULTS (WITH REFINED NAME MATCHING)\n", + "============================================================\n", + "\n", + "Total rows analyzed: 50,861\n", + "Rows with at least 2 models MEANINGFULLY agreeing (valid): 39,434 (77.5%)\n", + "Rows with all 3 models MEANINGFULLY agreeing: 21,459 (42.2%)\n", + "Rows with all-unknown values for any field: 2,010 (4.0%)\n", + "\n", + "Field-wise ANY agreement (including unknown matches):\n", + " - Country: 46,508 rows (91.4%)\n", + " - Gender: 49,884 rows (98.1%)\n", + " - Profession: 50,509 rows (99.3%)\n", + " - Name: 44,504 rows (87.5%)\n", + "\n", + "Field-wise MEANINGFUL agreement (excluding unknown matches):\n", + " - Country: 39,979 rows (78.6%)\n", + " - Gender: 49,197 rows (96.7%)\n", + " - Profession: 50,509 rows (99.3%)\n", + " - Name: 44,504 rows (87.5%)\n", + "\n", + "Examples of successful name matches with variations:\n", + " Row 0:\n", + " gemma: 'Lee Ji-eun' → cleaned: 'Lee Ji-eun'\n", + " qwen: 'Lee Ji-eun' → cleaned: 'Lee Ji-eun'\n", + " mistral: 'Lee Ji-eun (Lee, Ji-eun)' → cleaned: 'Lee Ji-eun'\n", + " Meaningful name agreements: 3/3\n", + " Row 2:\n", + " gemma: 'Li Yuchan' → cleaned: 'Li Yuchan'\n", + " qwen: 'Li Yuxiao' → cleaned: 'Li Yuxiao'\n", + " mistral: 'Liyuu (Full legal name unknown)' → cleaned: 'Liyuu'\n", + " Meaningful name agreements: 1/3\n", + " Row 3:\n", + " gemma: 'Irene Kim' → cleaned: 'Irene Kim'\n", + " qwen: 'Irene Kim' → cleaned: 'Irene Kim'\n", + " mistral: 'Irene (Full legal name unknown)' → cleaned: 'Irene'\n", + " Meaningful name agreements: 3/3\n", + "\n", + "Detailed MEANINGFUL Agreement Patterns:\n", + "\n", + "Number of MEANINGFUL agreeing pairs per field (out of 3 possible pairs):\n", + "\n", + "Country:\n", + " - 0 meaningful agreeing pairs: 10,882 rows (21.4%)\n", + " - 1 meaningful agreeing pairs: 16,159 rows (31.8%)\n", + " - 3 meaningful agreeing pairs: 23,820 rows (46.8%)\n", + "\n", + "Gender:\n", + " - 0 meaningful agreeing pairs: 1,664 rows (3.3%)\n", + " - 1 meaningful agreeing pairs: 5,190 rows (10.2%)\n", + " - 3 meaningful agreeing pairs: 44,007 rows (86.5%)\n", + "\n", + "Profession:\n", + " - 0 meaningful agreeing pairs: 352 rows (0.7%)\n", + " - 1 meaningful agreeing pairs: 3,916 rows (7.7%)\n", + " - 2 meaningful agreeing pairs: 3,239 rows (6.4%)\n", + " - 3 meaningful agreeing pairs: 43,354 rows (85.2%)\n", + "\n", + "Name:\n", + " - 0 meaningful agreeing pairs: 6,357 rows (12.5%)\n", + " - 1 meaningful agreeing pairs: 7,484 rows (14.7%)\n", + " - 2 meaningful agreeing pairs: 774 rows (1.5%)\n", + " - 3 meaningful agreeing pairs: 36,246 rows (71.3%)\n", + "\n", + "Adding consensus columns...\n", + "\n", + "Consensus column statistics:\n", + " - Consensus country: 48,970 rows (96.3%)\n", + " - Consensus gender: 50,775 rows (99.8%)\n", + " - Consensus primary profession: 50,849 rows (100.0%)\n", + "\n", + "Examples of primary profession consensus (first 5 rows with consensus):\n", + "\n", + " Row 0:\n", + " gemma: singer/musician (from singer/musician, actor, tv personality)\n", + " qwen: singer/musician (from singer/musician, model, public figure)\n", + " mistral: singer/musician (from singer/musician, tv personality, actress)\n", + " → Consensus: singer/musician\n", + "\n", + " Row 1:\n", + " gemma: model (from model, adult performer, online personality)\n", + " qwen: model (from model, online personality, artist)\n", + " mistral: model (from model, online personality, actress)\n", + " → Consensus: model\n", + "\n", + " Row 2:\n", + " gemma: singer/musician (from singer/musician, actor, model)\n", + " qwen: model (from model, online personality)\n", + " mistral: model (from model, online personality, artist)\n", + " → Consensus: model\n", + "\n", + " Row 3:\n", + " gemma: model (from model, online personality, actor)\n", + " qwen: singer/musician (from singer/musician, model, public figure)\n", + " mistral: actor (from actor, model, tv personality)\n", + " → Consensus: model\n", + "\n", + " Row 4:\n", + " gemma: singer/musician (from singer/musician, tv personality, online personality)\n", + " qwen: singer/musician (from singer/musician, model, online personality)\n", + " mistral: singer/musician (from singer/musician, tv personality, public figure)\n", + " → Consensus: singer/musician\n", + "\n", + "Saved analyzed data to: /shares/weddigen.ki.uzh/laura_wagner/phase_01/pm-paper/data/CSV/analyzed_llm_agreement.csv\n", + "Saved valid rows (39,434) to: /shares/weddigen.ki.uzh/laura_wagner/phase_01/pm-paper/data/CSV/analyzed_llm_agreement_valid.csv\n", + "Saved MEANINGFUL consensus rows (21,459) to: /shares/weddigen.ki.uzh/laura_wagner/phase_01/pm-paper/data/CSV/analyzed_llm_agreement_consensus.csv\n", + "Saved all-unknown rows (2,010) to: /shares/weddigen.ki.uzh/laura_wagner/phase_01/pm-paper/data/CSV/analyzed_llm_agreement_all_unknown.csv\n", + "\n", + "============================================================\n", + "SAMPLE OF MEANINGFUL CONSENSUS ROWS (first 2 rows)\n", + "============================================================\n", + "\n", + "Consensus Row 0:\n", + " Consensus Country: South Korea\n", + " Consensus Gender: Female\n", + " Consensus Primary Profession: singer/musician\n", + "\n", + " Individual Model Outputs:\n", + " gemma:\n", + " Name: Lee Ji-eun\n", + " Country: South Korea\n", + " Gender: Female\n", + " Profession: singer/musician, actor, tv personality\n", + " qwen:\n", + " Name: Lee Ji-eun\n", + " Country: South Korea\n", + " Gender: Female\n", + " Profession: singer/musician, model, public figure\n", + " mistral:\n", + " Name: Lee Ji-eun (Lee, Ji-eun)\n", + " Country: South Korea\n", + " Gender: Female\n", + " Profession: singer/musician, tv personality, actress\n", + "\n", + "Consensus Row 4:\n", + " Consensus Country: South Korea\n", + " Consensus Gender: Female\n", + " Consensus Primary Profession: singer/musician\n", + "\n", + " Individual Model Outputs:\n", + " gemma:\n", + " Name: Yoo Jimin\n", + " Country: South Korea\n", + " Gender: Female\n", + " Profession: singer/musician, tv personality, online personality\n", + " qwen:\n", + " Name: Kim Karina\n", + " Country: South Korea\n", + " Gender: Female\n", + " Profession: singer/musician, model, online personality\n", + " mistral:\n", + " Name: Karina (Kim Jung-yeon)\n", + " Country: South Korea\n", + " Gender: Female\n", + " Profession: singer/musician, tv personality, public figure\n" + ] + } + ], + "source": [ + "import pandas as pd\n", + "import numpy as np\n", + "import re\n", + "from typing import List, Set, Tuple\n", + "from collections import Counter\n", + "from pathlib import Path\n", + "\n", + "def is_unknown_value(value) -> bool:\n", + " \"\"\"Check if a value is considered 'unknown' or empty.\"\"\"\n", + " if pd.isna(value):\n", + " return True\n", + " \n", + " value_str = str(value).strip().lower()\n", + " \n", + " # Common unknown/empty indicators\n", + " unknown_patterns = [\n", + " 'unknown',\n", + " 'n/a',\n", + " 'na',\n", + " 'none',\n", + " '',\n", + " 'null',\n", + " 'not specified',\n", + " 'not available',\n", + " 'unspecified',\n", + " '?',\n", + " '??'\n", + " ]\n", + " \n", + " return value_str in unknown_patterns or len(value_str) == 0\n", + "\n", + "def clean_name_text(name: str) -> str:\n", + " \"\"\"Clean and normalize name text for comparison.\"\"\"\n", + " if pd.isna(name):\n", + " return \"\"\n", + " \n", + " name_str = str(name)\n", + " \n", + " # Remove content in parentheses, brackets, and their variations\n", + " name_str = re.sub(r'\\([^)]*\\)', '', name_str) # Remove (legal name unknown)\n", + " name_str = re.sub(r'\\[[^\\]]*\\]', '', name_str) # Remove [anything]\n", + " name_str = re.sub(r'\\{[^}]*\\}', '', name_str) # Remove {anything}\n", + " \n", + " # Remove common explanatory phrases\n", + " explanatory_phrases = [\n", + " r'legal name unknown',\n", + " r'full name',\n", + " r'also known as',\n", + " r'aka',\n", + " r'birth name',\n", + " r'real name',\n", + " r'stage name'\n", + " ]\n", + " \n", + " for phrase in explanatory_phrases:\n", + " name_str = re.sub(phrase, '', name_str, flags=re.IGNORECASE)\n", + " \n", + " # Normalize special characters\n", + " name_str = name_str.strip()\n", + " \n", + " # Replace multiple spaces with single space\n", + " name_str = re.sub(r'\\s+', ' ', name_str)\n", + " \n", + " return name_str\n", + "\n", + "def normalize_name(name: str) -> List[str]:\n", + " \"\"\"Normalize name by cleaning, converting to lowercase, and splitting into tokens.\"\"\"\n", + " if pd.isna(name):\n", + " return []\n", + " \n", + " # Clean the name text first\n", + " cleaned_name = clean_name_text(name)\n", + " \n", + " if not cleaned_name:\n", + " return []\n", + " \n", + " # Check if it's an unknown name after cleaning\n", + " if is_unknown_value(cleaned_name):\n", + " return []\n", + " \n", + " # Convert to lowercase\n", + " name_lower = cleaned_name.lower()\n", + " \n", + " # Normalize hyphenated names and special characters\n", + " # Convert \"na-eun\" to \"naeun\" for better matching\n", + " name_lower = re.sub(r'([a-z])-([a-z])', r'\\1\\2', name_lower)\n", + " \n", + " # Split into tokens\n", + " tokens = name_lower.split()\n", + " \n", + " # Remove common prefixes/suffixes that don't affect identity\n", + " common_elements = {'jr', 'sr', 'ii', 'iii', 'iv', 'v', 'dr', 'mr', 'mrs', 'ms', 'prof'}\n", + " tokens = [t for t in tokens if t not in common_elements]\n", + " \n", + " return tokens\n", + "\n", + "def names_match(name1: str, name2: str, threshold: float = 0.5) -> bool:\n", + " \"\"\"Check if two names match with refined partial matching.\"\"\"\n", + " # Check if either value is unknown\n", + " if is_unknown_value(name1) or is_unknown_value(name2):\n", + " return False\n", + " \n", + " tokens1 = normalize_name(name1)\n", + " tokens2 = normalize_name(name2)\n", + " \n", + " # If either is empty after normalization (was unknown), no match\n", + " if not tokens1 or not tokens2:\n", + " return False\n", + " \n", + " # Special case: if one name is clearly a subset of the other\n", + " # e.g., \"Olivia Holt\" vs \"Olivia Rose Holt\"\n", + " set1 = set(tokens1)\n", + " set2 = set(tokens2)\n", + " \n", + " # Check if one is subset of the other\n", + " if set1.issubset(set2) or set2.issubset(set1):\n", + " return True\n", + " \n", + " # Calculate token overlap\n", + " overlap = len(set1.intersection(set2))\n", + " \n", + " # Special handling for Asian names with/without hyphens\n", + " # Create versions without hyphens for comparison\n", + " tokens1_no_hyphen = [t.replace('-', '') for t in tokens1]\n", + " tokens2_no_hyphen = [t.replace('-', '') for t in tokens2]\n", + " \n", + " set1_no_hyphen = set(tokens1_no_hyphen)\n", + " set2_no_hyphen = set(tokens2_no_hyphen)\n", + " \n", + " overlap_no_hyphen = len(set1_no_hyphen.intersection(set2_no_hyphen))\n", + " \n", + " # Use the better overlap\n", + " best_overlap = max(overlap, overlap_no_hyphen)\n", + " \n", + " # Calculate match ratio based on the shorter name\n", + " min_length = min(len(tokens1), len(tokens2))\n", + " match_ratio = best_overlap / min_length if min_length > 0 else 0\n", + " \n", + " return match_ratio >= threshold\n", + "\n", + "def normalize_professions(prof_str: str) -> Set[str]:\n", + " \"\"\"Normalize profession strings to sets for comparison.\"\"\"\n", + " if pd.isna(prof_str) or is_unknown_value(prof_str):\n", + " return set()\n", + " \n", + " # Split by comma, strip whitespace, convert to lowercase\n", + " professions = [p.strip().lower() for p in str(prof_str).split(',') if p.strip() and not is_unknown_value(p)]\n", + " return set(professions)\n", + "\n", + "def professions_match(prof1: str, prof2: str) -> bool:\n", + " \"\"\"Check if two profession strings match (order doesn't matter).\"\"\"\n", + " set1 = normalize_professions(prof1)\n", + " set2 = normalize_professions(prof2)\n", + " \n", + " # Both empty or unknown - not a meaningful match\n", + " if not set1 and not set2:\n", + " return False\n", + " \n", + " # One is empty/unknown, the other has content - not a match\n", + " if (not set1 and set2) or (set1 and not set2):\n", + " return False\n", + " \n", + " # Match if they share at least one profession\n", + " return len(set1.intersection(set2)) > 0\n", + "\n", + "def check_model_agreement(row: pd.Series, models: List[str]) -> dict:\n", + " \"\"\"\n", + " Check agreement between models for a single row.\n", + " \n", + " Returns:\n", + " Dictionary with agreement results for this row\n", + " \"\"\"\n", + " results = {\n", + " 'row_valid': False,\n", + " 'all_agree': False,\n", + " 'agreements': {\n", + " 'country': 0,\n", + " 'gender': 0,\n", + " 'profession': 0,\n", + " 'name': 0\n", + " },\n", + " 'meaningful_agreements': {\n", + " 'country': 0, # Agreements that are NOT all unknown\n", + " 'gender': 0,\n", + " 'profession': 0,\n", + " 'name': 0\n", + " },\n", + " 'name_variations': [], # Store name variations for debugging\n", + " 'has_unknown_values': False # Track if row has unknown values\n", + " }\n", + " \n", + " # Collect values from each model\n", + " countries = [row.get(f'{model}_country', None) for model in models]\n", + " genders = [row.get(f'{model}_gender', None) for model in models]\n", + " professions = [row.get(f'{model}_profession_llm', None) for model in models]\n", + " raw_names = [row.get(f'{model}_full_name', None) for model in models]\n", + " cleaned_names = [clean_name_text(name) if pd.notna(name) else \"\" for name in raw_names]\n", + " \n", + " # Check if any field has all unknown values\n", + " all_countries_unknown = all(is_unknown_value(c) for c in countries)\n", + " all_genders_unknown = all(is_unknown_value(g) for g in genders)\n", + " all_professions_unknown = all(is_unknown_value(p) for p in professions)\n", + " all_names_unknown = all(is_unknown_value(n) for n in raw_names)\n", + " \n", + " results['has_unknown_values'] = (all_countries_unknown or all_genders_unknown or \n", + " all_professions_unknown or all_names_unknown)\n", + " \n", + " # Store name variations for debugging\n", + " results['name_variations'] = list(zip(models, raw_names, cleaned_names))\n", + " \n", + " # Count pairwise agreements\n", + " country_agreements = 0\n", + " gender_agreements = 0\n", + " profession_agreements = 0\n", + " name_agreements = 0\n", + " \n", + " # Count meaningful agreements (excluding unknown matches)\n", + " meaningful_country_agreements = 0\n", + " meaningful_gender_agreements = 0\n", + " meaningful_profession_agreements = 0\n", + " meaningful_name_agreements = 0\n", + " \n", + " # Check all pairs of models\n", + " n_models = len(models)\n", + " for i in range(n_models):\n", + " for j in range(i + 1, n_models):\n", + " # Country agreement (exact match)\n", + " if pd.notna(countries[i]) and pd.notna(countries[j]) and countries[i] == countries[j]:\n", + " country_agreements += 1\n", + " # Check if it's a meaningful agreement (not unknown)\n", + " if not is_unknown_value(countries[i]) and not is_unknown_value(countries[j]):\n", + " meaningful_country_agreements += 1\n", + " \n", + " # Gender agreement (exact match)\n", + " if pd.notna(genders[i]) and pd.notna(genders[j]) and genders[i] == genders[j]:\n", + " gender_agreements += 1\n", + " # Check if it's a meaningful agreement (not unknown)\n", + " if not is_unknown_value(genders[i]) and not is_unknown_value(genders[j]):\n", + " meaningful_gender_agreements += 1\n", + " \n", + " # Profession agreement (set-based with overlap)\n", + " if pd.notna(professions[i]) and pd.notna(professions[j]) and professions_match(professions[i], professions[j]):\n", + " profession_agreements += 1\n", + " # Check if it's a meaningful agreement (not unknown)\n", + " if not is_unknown_value(professions[i]) and not is_unknown_value(professions[j]):\n", + " meaningful_profession_agreements += 1\n", + " \n", + " # Name agreement (refined partial match)\n", + " if pd.notna(raw_names[i]) and pd.notna(raw_names[j]) and names_match(raw_names[i], raw_names[j]):\n", + " name_agreements += 1\n", + " # Check if it's a meaningful agreement (not unknown)\n", + " if not is_unknown_value(raw_names[i]) and not is_unknown_value(raw_names[j]):\n", + " meaningful_name_agreements += 1\n", + " \n", + " # Store agreement counts\n", + " results['agreements']['country'] = country_agreements\n", + " results['agreements']['gender'] = gender_agreements\n", + " results['agreements']['profession'] = profession_agreements\n", + " results['agreements']['name'] = name_agreements\n", + " \n", + " results['meaningful_agreements']['country'] = meaningful_country_agreements\n", + " results['meaningful_agreements']['gender'] = meaningful_gender_agreements\n", + " results['meaningful_agreements']['profession'] = meaningful_profession_agreements\n", + " results['meaningful_agreements']['name'] = meaningful_name_agreements\n", + " \n", + " # Check if row is valid (at least 2 models agree on all three fields)\n", + " # We need at least one MEANINGFUL agreement pair for country, gender, AND profession\n", + " valid_country = meaningful_country_agreements >= 1\n", + " valid_gender = meaningful_gender_agreements >= 1\n", + " valid_profession = meaningful_profession_agreements >= 1\n", + " \n", + " results['row_valid'] = valid_country and valid_gender and valid_profession\n", + " \n", + " # Check if all three models agree on everything MEANINGFULLY\n", + " # Need all possible pairs (3 pairs for 3 models) to have meaningful agreement\n", + " # AND we must have at least one non-unknown value per model for each field\n", + " all_agree_country = meaningful_country_agreements == 3 # 3 meaningful pairs\n", + " all_agree_gender = meaningful_gender_agreements == 3\n", + " all_agree_profession = meaningful_profession_agreements == 3\n", + " \n", + " # For names, we need all pairs to have meaningful agreement\n", + " all_agree_name = meaningful_name_agreements == 3\n", + " \n", + " # All three models must agree meaningfully on country, gender, and profession\n", + " results['all_agree'] = (all_agree_country and all_agree_gender and \n", + " all_agree_profession and not results['has_unknown_values'])\n", + " \n", + " return results\n", + "\n", + "def extract_consensus_values(row: pd.Series, models: List[str]) -> Tuple[str, str, str]:\n", + " \"\"\"\n", + " Extract consensus values for country, gender, and primary profession.\n", + " \n", + " Returns:\n", + " Tuple of (consensus_country, consensus_gender, consensus_primary_profession)\n", + " \"\"\"\n", + " # Extract country and gender (straightforward - they're the same in consensus rows)\n", + " countries = [row.get(f'{model}_country', None) for model in models]\n", + " genders = [row.get(f'{model}_gender', None) for model in models]\n", + " \n", + " # Get first non-null, non-unknown country and gender\n", + " consensus_country = None\n", + " for country in countries:\n", + " if pd.notna(country) and not is_unknown_value(country):\n", + " consensus_country = country\n", + " break\n", + " \n", + " consensus_gender = None\n", + " for gender in genders:\n", + " if pd.notna(gender) and not is_unknown_value(gender):\n", + " consensus_gender = gender\n", + " break\n", + " \n", + " # Extract primary profession (first profession from each model)\n", + " primary_professions = []\n", + " all_professions = [] # Track all professions for frequency counting\n", + " \n", + " for model in models:\n", + " prof_col = f'{model}_profession_llm'\n", + " prof_str = row.get(prof_col, None)\n", + " \n", + " if pd.notna(prof_str) and not is_unknown_value(prof_str):\n", + " # Split by comma and get professions\n", + " professions = [p.strip().lower() for p in str(prof_str).split(',') if p.strip()]\n", + " \n", + " if professions:\n", + " # First profession is the primary one\n", + " primary_professions.append(professions[0])\n", + " # Track all professions for frequency counting\n", + " all_professions.extend(professions)\n", + " \n", + " # Determine consensus primary profession\n", + " consensus_primary_profession = None\n", + " \n", + " if len(primary_professions) >= 2:\n", + " # Check if 2+ models agree on the same primary profession\n", + " primary_counter = Counter(primary_professions)\n", + " most_common_primary = primary_counter.most_common(1)[0]\n", + " \n", + " # If 2+ models agree on the primary profession, use it\n", + " if most_common_primary[1] >= 2:\n", + " consensus_primary_profession = most_common_primary[0]\n", + " else:\n", + " # All different - use the profession that appears most across ALL professions\n", + " all_counter = Counter(all_professions)\n", + " if all_counter:\n", + " consensus_primary_profession = all_counter.most_common(1)[0][0]\n", + " \n", + " return consensus_country, consensus_gender, consensus_primary_profession\n", + "\n", + "def add_consensus_columns(df_analyzed: pd.DataFrame, models: List[str]) -> pd.DataFrame:\n", + " \"\"\"\n", + " Add consensus columns for country, gender, and primary profession.\n", + " \"\"\"\n", + " print(\"\\nAdding consensus columns...\")\n", + " \n", + " consensus_data = []\n", + " for idx, row in df_analyzed.iterrows():\n", + " country, gender, profession = extract_consensus_values(row, models)\n", + " consensus_data.append({\n", + " 'consensus_country': country,\n", + " 'consensus_gender': gender,\n", + " 'consensus_primary_profession': profession\n", + " })\n", + " \n", + " consensus_df = pd.DataFrame(consensus_data)\n", + " \n", + " # Add columns to the analyzed dataframe\n", + " df_with_consensus = df_analyzed.copy()\n", + " df_with_consensus['consensus_country'] = consensus_df['consensus_country']\n", + " df_with_consensus['consensus_gender'] = consensus_df['consensus_gender']\n", + " df_with_consensus['consensus_primary_profession'] = consensus_df['consensus_primary_profession']\n", + " \n", + " # Print statistics about consensus values\n", + " total_rows = len(df_with_consensus)\n", + " \n", + " # Count non-null consensus values\n", + " country_count = df_with_consensus['consensus_country'].notna().sum()\n", + " gender_count = df_with_consensus['consensus_gender'].notna().sum()\n", + " profession_count = df_with_consensus['consensus_primary_profession'].notna().sum()\n", + " \n", + " print(f\"\\nConsensus column statistics:\")\n", + " print(f\" - Consensus country: {country_count:,} rows ({country_count/total_rows*100:.1f}%)\")\n", + " print(f\" - Consensus gender: {gender_count:,} rows ({gender_count/total_rows*100:.1f}%)\")\n", + " print(f\" - Consensus primary profession: {profession_count:,} rows ({profession_count/total_rows*100:.1f}%)\")\n", + " \n", + " # Show examples of profession consensus\n", + " print(\"\\nExamples of primary profession consensus (first 5 rows with consensus):\")\n", + " sample_rows = df_with_consensus[df_with_consensus['consensus_primary_profession'].notna()].head(5)\n", + " for idx, row in sample_rows.iterrows():\n", + " print(f\"\\n Row {idx}:\")\n", + " for model in models:\n", + " prof_col = f'{model}_profession_llm'\n", + " if prof_col in row and pd.notna(row[prof_col]):\n", + " profs = [p.strip() for p in str(row[prof_col]).split(',')]\n", + " print(f\" {model}: {profs[0] if profs else 'N/A'} (from {row[prof_col]})\")\n", + " print(f\" → Consensus: {row['consensus_primary_profession']}\")\n", + " \n", + " return df_with_consensus\n", + "\n", + "def analyze_model_agreement(df: pd.DataFrame, models: List[str]) -> pd.DataFrame:\n", + " \"\"\"\n", + " Analyze agreement between models and add agreement columns.\n", + " \n", + " Returns:\n", + " DataFrame with added agreement analysis columns\n", + " \"\"\"\n", + " print(f\"Analyzing agreement between models: {', '.join(models)}\")\n", + " print(f\"Total rows to analyze: {len(df)}\")\n", + " \n", + " # Prepare results storage\n", + " analysis_results = []\n", + " \n", + " # Track overall statistics\n", + " total_rows = len(df)\n", + " valid_rows = 0\n", + " all_agree_rows = 0\n", + " unknown_only_rows = 0 # Rows where all models say unknown for any field\n", + " field_agreements = {'country': 0, 'gender': 0, 'profession': 0, 'name': 0}\n", + " meaningful_field_agreements = {'country': 0, 'gender': 0, 'profession': 0, 'name': 0}\n", + " \n", + " # Track name matching examples for debugging\n", + " name_match_examples = []\n", + " name_mismatch_examples = []\n", + " \n", + " # Analyze each row\n", + " for idx, row in df.iterrows():\n", + " results = check_model_agreement(row, models)\n", + " analysis_results.append(results)\n", + " \n", + " # Update counters\n", + " if results['row_valid']:\n", + " valid_rows += 1\n", + " if results['all_agree']:\n", + " all_agree_rows += 1\n", + " if results['has_unknown_values']:\n", + " unknown_only_rows += 1\n", + " \n", + " # Track field agreements (rows with at least one agreement pair)\n", + " if results['agreements']['country'] > 0:\n", + " field_agreements['country'] += 1\n", + " if results['agreements']['gender'] > 0:\n", + " field_agreements['gender'] += 1\n", + " if results['agreements']['profession'] > 0:\n", + " field_agreements['profession'] += 1\n", + " if results['agreements']['name'] > 0:\n", + " field_agreements['name'] += 1\n", + " \n", + " # Track meaningful field agreements\n", + " if results['meaningful_agreements']['country'] > 0:\n", + " meaningful_field_agreements['country'] += 1\n", + " if results['meaningful_agreements']['gender'] > 0:\n", + " meaningful_field_agreements['gender'] += 1\n", + " if results['meaningful_agreements']['profession'] > 0:\n", + " meaningful_field_agreements['profession'] += 1\n", + " if results['meaningful_agreements']['name'] > 0:\n", + " meaningful_field_agreements['name'] += 1\n", + " \n", + " # Collect name examples for debugging\n", + " if len(name_match_examples) < 5 and results['meaningful_agreements']['name'] > 0:\n", + " name_variations = results['name_variations']\n", + " if any(n1 != n2 for _, n1, _ in name_variations for _, n2, _ in name_variations if n1 != n2):\n", + " name_match_examples.append({\n", + " 'row_idx': idx,\n", + " 'raw_names': [n for _, n, _ in name_variations],\n", + " 'cleaned_names': [c for _, _, c in name_variations],\n", + " 'agreements': results['meaningful_agreements']['name']\n", + " })\n", + " \n", + " # Convert results to DataFrame\n", + " results_df = pd.DataFrame(analysis_results)\n", + " \n", + " # Add agreement columns to original dataframe\n", + " df_analyzed = df.copy()\n", + " df_analyzed['row_valid'] = results_df['row_valid']\n", + " df_analyzed['all_models_agree'] = results_df['all_agree']\n", + " df_analyzed['has_unknown_values'] = results_df['has_unknown_values']\n", + " df_analyzed['country_agreements'] = results_df['agreements'].apply(lambda x: x['country'])\n", + " df_analyzed['gender_agreements'] = results_df['agreements'].apply(lambda x: x['gender'])\n", + " df_analyzed['profession_agreements'] = results_df['agreements'].apply(lambda x: x['profession'])\n", + " df_analyzed['name_agreements'] = results_df['agreements'].apply(lambda x: x['name'])\n", + " df_analyzed['meaningful_country_agreements'] = results_df['meaningful_agreements'].apply(lambda x: x['country'])\n", + " df_analyzed['meaningful_gender_agreements'] = results_df['meaningful_agreements'].apply(lambda x: x['gender'])\n", + " df_analyzed['meaningful_profession_agreements'] = results_df['meaningful_agreements'].apply(lambda x: x['profession'])\n", + " df_analyzed['meaningful_name_agreements'] = results_df['meaningful_agreements'].apply(lambda x: x['name'])\n", + " \n", + " # Calculate overall statistics\n", + " valid_percentage = (valid_rows / total_rows) * 100\n", + " all_agree_percentage = (all_agree_rows / total_rows) * 100\n", + " unknown_percentage = (unknown_only_rows / total_rows) * 100\n", + " \n", + " print(\"\\n\" + \"=\"*60)\n", + " print(\"AGREEMENT ANALYSIS RESULTS (WITH REFINED NAME MATCHING)\")\n", + " print(\"=\"*60)\n", + " print(f\"\\nTotal rows analyzed: {total_rows:,}\")\n", + " print(f\"Rows with at least 2 models MEANINGFULLY agreeing (valid): {valid_rows:,} ({valid_percentage:.1f}%)\")\n", + " print(f\"Rows with all 3 models MEANINGFULLY agreeing: {all_agree_rows:,} ({all_agree_percentage:.1f}%)\")\n", + " print(f\"Rows with all-unknown values for any field: {unknown_only_rows:,} ({unknown_percentage:.1f}%)\")\n", + " \n", + " print(f\"\\nField-wise ANY agreement (including unknown matches):\")\n", + " for field, count in field_agreements.items():\n", + " percentage = (count / total_rows) * 100\n", + " print(f\" - {field.capitalize()}: {count:,} rows ({percentage:.1f}%)\")\n", + " \n", + " print(f\"\\nField-wise MEANINGFUL agreement (excluding unknown matches):\")\n", + " for field, count in meaningful_field_agreements.items():\n", + " percentage = (count / total_rows) * 100\n", + " print(f\" - {field.capitalize()}: {count:,} rows ({percentage:.1f}%)\")\n", + " \n", + " # Show name matching examples\n", + " if name_match_examples:\n", + " print(f\"\\nExamples of successful name matches with variations:\")\n", + " for example in name_match_examples[:3]:\n", + " print(f\" Row {example['row_idx']}:\")\n", + " for model, raw_name, cleaned_name in zip(models, example['raw_names'], example['cleaned_names']):\n", + " print(f\" {model}: '{raw_name}' → cleaned: '{cleaned_name}'\")\n", + " print(f\" Meaningful name agreements: {example['agreements']}/3\")\n", + " \n", + " # Detailed breakdown of agreement patterns\n", + " print(f\"\\nDetailed MEANINGFUL Agreement Patterns:\")\n", + " \n", + " # Count different meaningful agreement levels\n", + " print(f\"\\nNumber of MEANINGFUL agreeing pairs per field (out of 3 possible pairs):\")\n", + " for field in ['country', 'gender', 'profession', 'name']:\n", + " col_name = f'meaningful_{field}_agreements'\n", + " print(f\"\\n{field.capitalize()}:\")\n", + " for i in range(4): # 0, 1, 2, or 3 agreeing pairs\n", + " count = (df_analyzed[col_name] == i).sum()\n", + " if count > 0:\n", + " percentage = (count / total_rows) * 100\n", + " print(f\" - {i} meaningful agreeing pairs: {count:,} rows ({percentage:.1f}%)\")\n", + " \n", + " return df_analyzed\n", + "\n", + "def save_analysis_results(df_analyzed: pd.DataFrame, output_path: Path):\n", + " \"\"\"Save the analyzed dataframe with agreement columns.\"\"\"\n", + " # Save full analyzed dataset\n", + " df_analyzed.to_csv(output_path, index=False)\n", + " print(f\"\\nSaved analyzed data to: {output_path}\")\n", + " \n", + " # Save only valid rows (at least 2 models meaningfully agree)\n", + " valid_path = output_path.parent / f\"{output_path.stem}_valid{output_path.suffix}\"\n", + " valid_df = df_analyzed[df_analyzed['row_valid']]\n", + " valid_df.to_csv(valid_path, index=False)\n", + " print(f\"Saved valid rows ({len(valid_df):,}) to: {valid_path}\")\n", + " \n", + " # Save rows where all 3 models MEANINGFULLY agree (excluding unknown agreements)\n", + " consensus_path = output_path.parent / f\"{output_path.stem}_consensus{output_path.suffix}\"\n", + " consensus_df = df_analyzed[df_analyzed['all_models_agree']]\n", + " consensus_df.to_csv(consensus_path, index=False)\n", + " print(f\"Saved MEANINGFUL consensus rows ({len(consensus_df):,}) to: {consensus_path}\")\n", + " \n", + " # Also save rows that were excluded due to all-unknown values\n", + " unknown_path = output_path.parent / f\"{output_path.stem}_all_unknown{output_path.suffix}\"\n", + " unknown_df = df_analyzed[df_analyzed['has_unknown_values']]\n", + " unknown_df.to_csv(unknown_path, index=False)\n", + " print(f\"Saved all-unknown rows ({len(unknown_df):,}) to: {unknown_path}\")\n", + "\n", + "# ============================================================\n", + "# EXAMPLE USAGE\n", + "# ============================================================\n", + "\n", + "if __name__ == \"__main__\":\n", + " # Set up paths (adjust to your directory structure)\n", + " current_dir = path.cwd()\n", + " \n", + " # Load the combined data\n", + " combined_file = current_dir.parent / \"data/CSV/combined_llm_annotations.csv\"\n", + " print(f\"Loading combined data from: {combined_file}\")\n", + " \n", + " if combined_file.exists():\n", + " df_combined = pd.read_csv(combined_file)\n", + " print(f\"Loaded {len(df_combined):,} rows with {len(df_combined.columns)} columns\")\n", + " \n", + " # Define the models\n", + " models = ['gemma', 'qwen', 'mistral']\n", + " \n", + " # Run the agreement analysis\n", + " df_analyzed = analyze_model_agreement(df_combined, models)\n", + " \n", + " # Add consensus columns\n", + " df_with_consensus = add_consensus_columns(df_analyzed, models)\n", + " \n", + " # Save the results\n", + " output_file = current_dir.parent / \"data/CSV/analyzed_llm_agreement.csv\"\n", + " save_analysis_results(df_with_consensus, output_file)\n", + " \n", + " # Show sample of consensus rows\n", + " print(\"\\n\" + \"=\"*60)\n", + " print(\"SAMPLE OF MEANINGFUL CONSENSUS ROWS (first 2 rows)\")\n", + " print(\"=\"*60)\n", + " \n", + " consensus_rows = df_with_consensus[df_with_consensus['all_models_agree']].head(2)\n", + " \n", + " if len(consensus_rows) > 0:\n", + " for idx, row in consensus_rows.iterrows():\n", + " print(f\"\\nConsensus Row {idx}:\")\n", + " print(f\" Consensus Country: {row['consensus_country']}\")\n", + " print(f\" Consensus Gender: {row['consensus_gender']}\")\n", + " print(f\" Consensus Primary Profession: {row['consensus_primary_profession']}\")\n", + " print(f\"\\n Individual Model Outputs:\")\n", + " for model in models:\n", + " name_col = f'{model}_full_name'\n", + " country_col = f'{model}_country'\n", + " gender_col = f'{model}_gender'\n", + " prof_col = f'{model}_profession_llm'\n", + " \n", + " if name_col in row:\n", + " print(f\" {model}:\")\n", + " print(f\" Name: {row[name_col]}\")\n", + " print(f\" Country: {row[country_col]}\")\n", + " print(f\" Gender: {row[gender_col]}\")\n", + " print(f\" Profession: {row[prof_col]}\")\n", + " else:\n", + " print(f\"Error: Could not find combined data file at {combined_file}\")\n", + " print(\"Please adjust the path in the script to point to your data file.\")" + ] + }, + { + "cell_type": "markdown", + "id": "2d761d7f", + "metadata": {}, + "source": [ + "# Improved Consensus Script with Position-Aware Logic\n", + "# Specifically addresses the \"adult performer\" underrepresentation problem" + ] + }, + { + "cell_type": "code", + "execution_count": 2, + "id": "2aac8386", + "metadata": {}, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "================================================================================\n", + "IMPROVED CONSENSUS CREATION\n", + "================================================================================\n", + "Method: hybrid\n", + "Reading: /home/lauhp/000_PHD/000_010_PUBLICATION/CODE/pm-paper/data/CSV/combined_llm_annotations.csv\n", + "Input shape: (50861, 67)\n", + "Models: gemma, mistral, qwen\n", + "\n", + "Processing rows...\n", + " Processed 50000/50861 rows...\n", + " Processed 50861 rows. \n", + "\n", + "================================================================================\n", + "RESULTS\n", + "================================================================================\n", + "Total input rows: 50,861\n", + "Rows passing all criteria: 23,084 (45.4%)\n", + "\n", + "Consensus method usage:\n", + " - weighted: 43,445 (188.2%)\n", + " - adult_special: 5,622 (24.4%)\n", + " - adult_any_position: 1,784 (7.7%)\n", + " - no_consensus: 10 (0.0%)\n", + "\n", + "================================================================================\n", + "PROFESSION DISTRIBUTION\n", + "================================================================================\n", + "\n", + "Top 20 professions:\n", + "consensus_profession\n", + "actor 11046\n", + "singer/musician 3939\n", + "model 3017\n", + "online personality 1651\n", + "adult performer 1206\n", + "public figure 986\n", + "sports professional 458\n", + "voice actor/asmr 376\n", + "tv personality 373\n", + "wrestler 10\n", + "comedian 8\n", + "cheerleader 2\n", + "actress 2\n", + "dancer 2\n", + "architect 1\n", + "entrepreneur 1\n", + "basketball player 1\n", + "podcaster 1\n", + "artist 1\n", + "gymnast 1\n", + "Name: count, dtype: int64\n", + "\n", + "🎯 Adult performer variants: 1206 (5.22%)\n", + "\n", + "================================================================================\n", + "✓ Improved consensus saved to: improved_consensus.csv\n", + " Total rows: 23,084\n", + "================================================================================\n", + "\n", + "✅ Complete!\n" + ] + } + ], + "source": [ + "\n", + "\n", + "import pandas as pd\n", + "from pathlib import Path\n", + "from collections import Counter\n", + "import re\n", + "\n", + "# ============================================================\n", + "# CONFIGURATION\n", + "# ============================================================\n", + "\n", + "# Semantic grouping for professions\n", + "PROFESSION_GROUPS = {\n", + " 'adult_entertainment': [\n", + " 'adult performer', 'adult film', 'pornstar', 'av actress', \n", + " 'av idol', 'jav idol', 'adult model', 'adult entertainer'\n", + " ],\n", + " 'mainstream_model': [\n", + " 'model', 'fashion model', 'instagram model', 'supermodel',\n", + " 'runway model', 'commercial model'\n", + " ],\n", + " 'actor': [\n", + " 'actor', 'actress', 'film actor', 'tv actor', 'television actor'\n", + " ],\n", + " 'musician': [\n", + " 'singer', 'musician', 'singer/musician', 'music artist', 'vocalist'\n", + " ],\n", + " 'online_personality': [\n", + " 'online personality', 'influencer', 'content creator', \n", + " 'youtuber', 'streamer', 'social media personality'\n", + " ],\n", + " 'tv_personality': [\n", + " 'tv personality', 'television personality', 'tv host', 'presenter'\n", + " ]\n", + "}\n", + "\n", + "# Position weights (1st mention = most important)\n", + "POSITION_WEIGHTS = {\n", + " 0: 3.0, # First position\n", + " 1: 2.0, # Second position\n", + " 2: 1.0 # Third position\n", + "}\n", + "\n", + "# Model reliability weights (based on analysis)\n", + "MODEL_WEIGHTS = {\n", + " 'gemma': 1.0,\n", + " 'mistral': 0.85, # Slightly lower due to 76% detection rate vs 92%\n", + " 'qwen': 1.0\n", + "}\n", + "\n", + "# ============================================================\n", + "# UTILITY FUNCTIONS\n", + "# ============================================================\n", + "\n", + "def normalize_value(value):\n", + " \"\"\"Normalize values for comparison (handle NaN, whitespace, case)\"\"\"\n", + " if pd.isna(value):\n", + " return None\n", + " return str(value).strip().lower()\n", + "\n", + "def is_unknown_value(value):\n", + " \"\"\"Check if a value represents 'unknown' or similar non-informative values\"\"\"\n", + " if value is None:\n", + " return True\n", + " \n", + " value_str = str(value).strip().lower()\n", + " \n", + " unknown_patterns = [\n", + " 'unknown', 'n/a', 'na', 'none', 'not specified', \n", + " 'not available', 'unclear', 'uncertain', '', 'null'\n", + " ]\n", + " \n", + " return value_str in unknown_patterns\n", + "\n", + "def parse_profession_list(profession_str):\n", + " \"\"\"Parse comma-separated profession list into normalized list\"\"\"\n", + " if pd.isna(profession_str):\n", + " return []\n", + " \n", + " professions = [p.strip().lower() for p in str(profession_str).split(',')]\n", + " return [p for p in professions if p and not is_unknown_value(p)]\n", + "\n", + "def find_profession_group(profession, groups=PROFESSION_GROUPS):\n", + " \"\"\"Find which semantic group a profession belongs to\"\"\"\n", + " profession_lower = profession.lower()\n", + " \n", + " for group_name, terms in groups.items():\n", + " if profession_lower in terms:\n", + " return group_name\n", + " # Partial match for compound terms\n", + " if any(term in profession_lower for term in terms):\n", + " return group_name\n", + " \n", + " return profession_lower # Return as-is if no group found\n", + "\n", + "# ============================================================\n", + "# CONSENSUS ALGORITHMS\n", + "# ============================================================\n", + "\n", + "def get_position_weighted_consensus(profession_lists, model_names, \n", + " weights=POSITION_WEIGHTS, \n", + " model_weights=MODEL_WEIGHTS):\n", + " \"\"\"\n", + " Get consensus using position-based weighting.\n", + " \n", + " Professions mentioned first get more weight than those mentioned second or third.\n", + " Different models can have different reliability weights.\n", + " \n", + " Returns: (consensus_profession, weighted_score, breakdown_dict)\n", + " \"\"\"\n", + " scores = {}\n", + " breakdown = {}\n", + " \n", + " for model, prof_list in zip(model_names, profession_lists):\n", + " professions = parse_profession_list(prof_list)\n", + " model_weight = model_weights.get(model, 1.0)\n", + " \n", + " for position, profession in enumerate(professions[:3]): # Only top 3\n", + " pos_weight = weights.get(position, 0)\n", + " score = pos_weight * model_weight\n", + " \n", + " scores[profession] = scores.get(profession, 0) + score\n", + " \n", + " if profession not in breakdown:\n", + " breakdown[profession] = []\n", + " breakdown[profession].append({\n", + " 'model': model,\n", + " 'position': position + 1,\n", + " 'weight': score\n", + " })\n", + " \n", + " if not scores:\n", + " return None, 0, {}\n", + " \n", + " best_profession = max(scores.items(), key=lambda x: x[1])\n", + " return best_profession[0], best_profession[1], breakdown\n", + "\n", + "def get_semantic_consensus(profession_lists, model_names, required_agreement=2):\n", + " \"\"\"\n", + " Get consensus by grouping semantically similar professions.\n", + " \n", + " First votes on profession categories (e.g., \"adult entertainment\"),\n", + " then picks the most common specific term within the winning category.\n", + " \n", + " Returns: (consensus_profession, agreement_count, category)\n", + " \"\"\"\n", + " category_votes = {}\n", + " profession_within_category = {}\n", + " \n", + " for model, prof_list in zip(model_names, profession_lists):\n", + " professions = parse_profession_list(prof_list)\n", + " \n", + " if not professions:\n", + " continue\n", + " \n", + " # Use first profession from each model for category voting\n", + " first_prof = professions[0]\n", + " category = find_profession_group(first_prof)\n", + " \n", + " category_votes[category] = category_votes.get(category, 0) + 1\n", + " \n", + " if category not in profession_within_category:\n", + " profession_within_category[category] = []\n", + " profession_within_category[category].append(first_prof)\n", + " \n", + " if not category_votes:\n", + " return None, 0, None\n", + " \n", + " # Find winning category\n", + " winning_category, count = max(category_votes.items(), key=lambda x: x[1])\n", + " \n", + " if count < required_agreement:\n", + " return None, count, winning_category\n", + " \n", + " # Pick most common specific term within winning category\n", + " specific_terms = profession_within_category[winning_category]\n", + " most_common_term = Counter(specific_terms).most_common(1)[0][0]\n", + " \n", + " return most_common_term, count, winning_category\n", + "\n", + "def get_any_position_consensus(profession_lists, target_profession, \n", + " required_agreement=2):\n", + " \"\"\"\n", + " Check if target profession appears ANYWHERE in the lists.\n", + " \n", + " Useful for professions that are consistently mentioned but not always first.\n", + " \n", + " Returns: (found, agreement_count, positions_found)\n", + " \"\"\"\n", + " count = 0\n", + " positions = []\n", + " \n", + " for prof_list in profession_lists:\n", + " professions = parse_profession_list(prof_list)\n", + " \n", + " for i, prof in enumerate(professions):\n", + " if target_profession.lower() in prof.lower():\n", + " count += 1\n", + " positions.append(i + 1)\n", + " break\n", + " \n", + " return count >= required_agreement, count, positions\n", + "\n", + "def get_hybrid_consensus(profession_lists, model_names):\n", + " \"\"\"\n", + " Hybrid consensus strategy that tries multiple approaches.\n", + " \n", + " Strategy priority:\n", + " 1. Check for \"adult performer\" anywhere in lists (special case)\n", + " 2. Position-weighted consensus\n", + " 3. Semantic consensus\n", + " 4. Fallback to most common first profession\n", + " \n", + " Returns: (consensus_profession, method_used, confidence_score)\n", + " \"\"\"\n", + " # Special case: Check for adult entertainment professions\n", + " adult_terms = PROFESSION_GROUPS['adult_entertainment']\n", + " for term in adult_terms:\n", + " found, count, positions = get_any_position_consensus(profession_lists, term, required_agreement=2)\n", + " if found:\n", + " # Use weighted consensus to pick the exact term\n", + " weighted_prof, score, _ = get_position_weighted_consensus(profession_lists, model_names)\n", + " \n", + " # Check if the weighted winner is an adult entertainment term\n", + " if any(term in weighted_prof for term in adult_terms):\n", + " return weighted_prof, 'adult_special', score\n", + " \n", + " # If weighted winner isn't adult term, but 2+ models mentioned it, use it\n", + " return term, 'adult_any_position', count * 2.0\n", + " \n", + " # Try position-weighted consensus\n", + " weighted_prof, score, breakdown = get_position_weighted_consensus(profession_lists, model_names)\n", + " \n", + " if score >= 3.0: # Reasonable threshold (e.g., 2 models first position)\n", + " return weighted_prof, 'weighted', score\n", + " \n", + " # Try semantic consensus\n", + " semantic_prof, count, category = get_semantic_consensus(profession_lists, model_names)\n", + " \n", + " if count >= 2:\n", + " return semantic_prof, 'semantic', count * 1.5\n", + " \n", + " # Fallback: most common first profession\n", + " first_professions = []\n", + " for prof_list in profession_lists:\n", + " professions = parse_profession_list(prof_list)\n", + " if professions:\n", + " first_professions.append(professions[0])\n", + " \n", + " if first_professions:\n", + " most_common = Counter(first_professions).most_common(1)[0]\n", + " if most_common[1] >= 2:\n", + " return most_common[0], 'simple_majority', most_common[1]\n", + " \n", + " return None, 'no_consensus', 0\n", + "\n", + "# ============================================================\n", + "# ORIGINAL CONSENSUS (for comparison)\n", + "# ============================================================\n", + "\n", + "def get_consensus_value_original(values, required_agreement=2):\n", + " \"\"\"Original consensus method - compares entire strings\"\"\"\n", + " normalized = [normalize_value(v) for v in values]\n", + " valid_values = [v for v in normalized if v is not None]\n", + " \n", + " if not valid_values:\n", + " return None, 0\n", + " \n", + " value_counts = Counter(valid_values)\n", + " most_common_value, count = value_counts.most_common(1)[0]\n", + " \n", + " if count >= required_agreement:\n", + " return most_common_value, count\n", + " else:\n", + " return None, count\n", + "\n", + "# ============================================================\n", + "# MAIN CONSENSUS CREATION\n", + "# ============================================================\n", + "\n", + "def create_improved_consensus(input_file, output_file, \n", + " models=['gemma', 'mistral', 'qwen'],\n", + " consensus_method='hybrid'):\n", + " \"\"\"\n", + " Create improved consensus CSV with better profession detection.\n", + " \n", + " Parameters:\n", + " - input_file: Path to combined_llm_annotations.csv\n", + " - output_file: Path to save improved consensus\n", + " - models: List of model names\n", + " - consensus_method: 'hybrid', 'weighted', 'semantic', or 'original'\n", + " \"\"\"\n", + " print(\"=\"*80)\n", + " print(\"IMPROVED CONSENSUS CREATION\")\n", + " print(\"=\"*80)\n", + " print(f\"Method: {consensus_method}\")\n", + " print(f\"Reading: {input_file}\")\n", + " \n", + " df = pd.read_csv(input_file)\n", + " \n", + " print(f\"Input shape: {df.shape}\")\n", + " print(f\"Models: {', '.join(models)}\")\n", + " \n", + " consensus_data = []\n", + " stats = {\n", + " 'total_rows': len(df),\n", + " 'country_fail': 0,\n", + " 'gender_fail': 0,\n", + " 'profession_fail': 0,\n", + " 'unknown_values': 0,\n", + " 'all_pass': 0,\n", + " 'method_counts': {}\n", + " }\n", + " \n", + " print(\"\\nProcessing rows...\")\n", + " \n", + " for idx, row in df.iterrows():\n", + " if idx % 1000 == 0:\n", + " print(f\" Processed {idx}/{len(df)} rows...\", end='\\r')\n", + " \n", + " # Get values for each field\n", + " countries = [row[f'{model}_country'] for model in models]\n", + " genders = [row[f'{model}_gender'] for model in models]\n", + " professions = [row[f'{model}_profession_llm'] for model in models]\n", + " \n", + " # Country consensus (strict: all 3 must agree)\n", + " country_consensus, country_count = get_consensus_value_original(countries, required_agreement=3)\n", + " \n", + " # Gender consensus (strict: all 3 must agree)\n", + " gender_consensus, gender_count = get_consensus_value_original(genders, required_agreement=3)\n", + " \n", + " # Profession consensus (IMPROVED)\n", + " if consensus_method == 'hybrid':\n", + " profession_consensus, method, prof_score = get_hybrid_consensus(professions, models)\n", + " prof_count = int(prof_score / 1.5) # Rough conversion to count\n", + " elif consensus_method == 'weighted':\n", + " profession_consensus, prof_score, _ = get_position_weighted_consensus(professions, models)\n", + " prof_count = int(prof_score / 2)\n", + " method = 'weighted'\n", + " elif consensus_method == 'semantic':\n", + " profession_consensus, prof_count, _ = get_semantic_consensus(professions, models)\n", + " method = 'semantic'\n", + " prof_score = prof_count\n", + " else: # original\n", + " # Get first profession from each model\n", + " first_profs = [parse_profession_list(p)[0] if parse_profession_list(p) else None \n", + " for p in professions]\n", + " profession_consensus, prof_count = get_consensus_value_original(first_profs, required_agreement=2)\n", + " method = 'original'\n", + " prof_score = prof_count\n", + " \n", + " # Track method usage\n", + " stats['method_counts'][method] = stats['method_counts'].get(method, 0) + 1\n", + " \n", + " # Determine if row passes\n", + " country_pass = country_count == 3\n", + " gender_pass = gender_count == 3\n", + " profession_pass = profession_consensus is not None\n", + " \n", + " has_unknown = (\n", + " is_unknown_value(country_consensus) or \n", + " is_unknown_value(gender_consensus) or \n", + " is_unknown_value(profession_consensus)\n", + " )\n", + " \n", + " if not country_pass:\n", + " stats['country_fail'] += 1\n", + " if not gender_pass:\n", + " stats['gender_fail'] += 1\n", + " if not profession_pass:\n", + " stats['profession_fail'] += 1\n", + " if has_unknown:\n", + " stats['unknown_values'] += 1\n", + " \n", + " if country_pass and gender_pass and profession_pass and not has_unknown:\n", + " stats['all_pass'] += 1\n", + " \n", + " consensus_data.append({\n", + " 'row_index': idx,\n", + " 'consensus_country': country_consensus,\n", + " 'consensus_gender': gender_consensus,\n", + " 'consensus_profession': profession_consensus,\n", + " 'profession_method': method,\n", + " 'profession_confidence': prof_score\n", + " })\n", + " \n", + " print(f\"\\n Processed {len(df)} rows. \")\n", + " \n", + " # Create result dataframe\n", + " if consensus_data:\n", + " result_df = df.iloc[[c['row_index'] for c in consensus_data]].copy().reset_index(drop=True)\n", + " \n", + " # Add consensus columns\n", + " for key in ['consensus_country', 'consensus_gender', 'consensus_profession', \n", + " 'profession_method', 'profession_confidence']:\n", + " result_df[key] = [c[key] for c in consensus_data]\n", + " \n", + " # Reorder columns (consensus columns first)\n", + " consensus_cols = ['consensus_country', 'consensus_gender', 'consensus_profession',\n", + " 'profession_method', 'profession_confidence']\n", + " other_cols = [c for c in result_df.columns if c not in consensus_cols]\n", + " result_df = result_df[consensus_cols + other_cols]\n", + " \n", + " # Save\n", + " result_df.to_csv(output_file, index=False)\n", + " \n", + " print(\"\\n\" + \"=\"*80)\n", + " print(\"RESULTS\")\n", + " print(\"=\"*80)\n", + " print(f\"Total input rows: {stats['total_rows']:,}\")\n", + " print(f\"Rows passing all criteria: {stats['all_pass']:,} ({stats['all_pass']/stats['total_rows']*100:.1f}%)\")\n", + " \n", + " print(f\"\\nConsensus method usage:\")\n", + " for method, count in sorted(stats['method_counts'].items(), key=lambda x: -x[1]):\n", + " print(f\" - {method}: {count:,} ({count/stats['all_pass']*100:.1f}%)\")\n", + " \n", + " print(\"\\n\" + \"=\"*80)\n", + " print(\"PROFESSION DISTRIBUTION\")\n", + " print(\"=\"*80)\n", + " print(\"\\nTop 20 professions:\")\n", + " print(result_df['consensus_profession'].value_counts().head(20))\n", + " \n", + " # Specifically check adult performer\n", + " adult_count = result_df['consensus_profession'].apply(\n", + " lambda x: 'adult' in str(x).lower() if pd.notna(x) else False\n", + " ).sum()\n", + " print(f\"\\n🎯 Adult performer variants: {adult_count} ({adult_count/len(result_df)*100:.2f}%)\")\n", + " \n", + " print(\"\\n\" + \"=\"*80)\n", + " print(f\"✓ Improved consensus saved to: {output_file.name}\")\n", + " print(f\" Total rows: {len(result_df):,}\")\n", + " print(\"=\"*80)\n", + " \n", + " return result_df\n", + " else:\n", + " print(\"\\n⚠ WARNING: No rows passed all criteria!\")\n", + " return pd.DataFrame()\n", + "\n", + "# ============================================================\n", + "# COMPARISON FUNCTION\n", + "# ============================================================\n", + "\n", + "def compare_consensus_methods(input_file, models=['gemma', 'mistral', 'qwen']):\n", + " \"\"\"Compare different consensus methods side by side\"\"\"\n", + " \n", + " print(\"=\"*80)\n", + " print(\"CONSENSUS METHOD COMPARISON\")\n", + " print(\"=\"*80)\n", + " \n", + " methods = ['original', 'weighted', 'semantic', 'hybrid']\n", + " results = {}\n", + " \n", + " for method in methods:\n", + " print(f\"\\n--- Testing {method} method ---\")\n", + " \n", + " output_file = Path(input_file).parent / f\"consensus_{method}.csv\"\n", + " result_df = create_improved_consensus(input_file, output_file, models, method)\n", + " \n", + " if len(result_df) > 0:\n", + " adult_count = result_df['consensus_profession'].apply(\n", + " lambda x: 'adult' in str(x).lower() if pd.notna(x) else False\n", + " ).sum()\n", + " \n", + " results[method] = {\n", + " 'total_rows': len(result_df),\n", + " 'adult_performer_count': adult_count,\n", + " 'adult_performer_pct': adult_count / len(result_df) * 100\n", + " }\n", + " \n", + " print(\"\\n\" + \"=\"*80)\n", + " print(\"COMPARISON SUMMARY\")\n", + " print(\"=\"*80)\n", + " \n", + " print(f\"\\n{'Method':<15} {'Total Rows':<12} {'Adult Performer':<16} {'% Adult':<10}\")\n", + " print(\"-\" * 65)\n", + " \n", + " for method, stats in results.items():\n", + " print(f\"{method:<15} {stats['total_rows']:<12,} {stats['adult_performer_count']:<16,} {stats['adult_performer_pct']:<10.2f}%\")\n", + " \n", + " if 'original' in results and 'hybrid' in results:\n", + " improvement = results['hybrid']['adult_performer_count'] - results['original']['adult_performer_count']\n", + " pct_improvement = improvement / results['original']['adult_performer_count'] * 100\n", + " \n", + " print(f\"\\n✨ Hybrid method improvement over original:\")\n", + " print(f\" +{improvement} adult performer cases (+{pct_improvement:.1f}%)\")\n", + "\n", + "# ============================================================\n", + "# MAIN EXECUTION\n", + "# ============================================================\n", + "\n", + "if __name__ == \"__main__\":\n", + " current_dir = Path.cwd()\n", + " \n", + " input_file = current_dir.parent / \"data/CSV/combined_llm_annotations.csv\"\n", + " output_file = current_dir.parent / \"data/CSV/improved_consensus.csv\"\n", + " \n", + " if not input_file.exists():\n", + " print(f\"Error: Input file not found: {input_file}\")\n", + " else:\n", + " # Run comparison (comment out if you just want hybrid)\n", + " # compare_consensus_methods(input_file)\n", + " \n", + " # Or run single method (hybrid recommended)\n", + " result_df = create_improved_consensus(\n", + " input_file, \n", + " output_file, \n", + " models=['gemma', 'mistral', 'qwen'],\n", + " consensus_method='hybrid'\n", + " )\n", + " \n", + " print(\"\\n✅ Complete!\")" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "70403f7c-6ea9-4f21-9704-aec0c37a591b", + "metadata": {}, + "outputs": [], + "source": [] + } + ], + "metadata": { + "kernelspec": { + "display_name": "latm", + "language": "python", + "name": "python3" + }, + "language_info": { + "codemirror_mode": { + "name": "ipython", + "version": 3 + }, + "file_extension": ".py", + "mimetype": "text/x-python", + "name": "python", + "nbconvert_exporter": "python", + "pygments_lexer": "ipython3", + "version": "3.10.15" + } + }, + "nbformat": 4, + "nbformat_minor": 5 +} diff --git a/jupyter_notebooks/.ipynb_checkpoints/Section_2-4_Figure_9_ectract_LoRA_metadata_v2-checkpoint.ipynb b/jupyter_notebooks/.ipynb_checkpoints/Section_2-4_Figure_9_ectract_LoRA_metadata_v2-checkpoint.ipynb new file mode 100644 index 0000000000000000000000000000000000000000..7ad7463163f1d7a7ffdc81a2de86f326c47c551d --- /dev/null +++ b/jupyter_notebooks/.ipynb_checkpoints/Section_2-4_Figure_9_ectract_LoRA_metadata_v2-checkpoint.ipynb @@ -0,0 +1,400 @@ +{ + "cells": [ + { + "cell_type": "markdown", + "id": "f36422c8", + "metadata": {}, + "source": [ + "# LoRA metadata" + ] + }, + { + "cell_type": "raw", + "id": "8a2feb6e", + "metadata": {}, + "source": [ + "LoRA Metadata Processing Workflow\n", + "┌──────────────────────┐ ┌───────────────────────────────────────────────────┐\n", + "│ Load CSV File │ --> │ Read adapter metadata CSV file. │\n", + "│ Read Model Versions │ │ Extract model version IDs and relevant data. │\n", + "└─────────┬────────────┘ └───────────────────────────────────────────────────┘\n", + " ↓\n", + "┌──────────────────────┐ ┌───────────────────────────────────────────────────┐\n", + "│ Download Adapter │ --> │ Use stored download URLs to fetch adapter files │\n", + "│ Files Using API │ │ using rotating API keys. │\n", + "└─────────┬────────────┘ └───────────────────────────────────────────────────┘\n", + "┌──────────────────────┐ ┌───────────────────────────────────────────────────┐\n", + "│ Parse Metadata │ --> │ Extract safetensors metadata, such as training │\n", + "│ from SafeTensor │ │ images, model type, and architecture. │\n", + "│ Files │ │ │\n", + "└─────────┬────────────┘ └───────────────────────────────────────────────────┘\n", + " ↓\n", + "┌──────────────────────┐ ┌───────────────────────────────────────────────────┐\n", + "│ Store Parsed │ --> │ Save extracted metadata into structured JSON │\n", + "│ Metadata as JSON │ │ files for later analysis. │\n", + "└─────────┬────────────┘ └───────────────────────────────────────────────────┘\n", + "\n", + "┌──────────────────────┐ ┌───────────────────────────────────────────────────┐\n", + "│ Process JSON Files │ --> │ Read saved JSON metadata, extract relevant │\n", + "│ for Consolidation │ │ details, and filter necessary attributes. │\n", + "└─────────┬────────────┘ └───────────────────────────────────────────────────┘\n", + " ↓\n", + "┌──────────────────────┐ ┌───────────────────────────────────────────────────┐\n", + "│ Extract Training │ --> │ Identify most frequent training tags, architectures│\n", + "│ Tags & Model Info │ │ and systems used for model creation. │\n", + "└─────────┬────────────┘ └───────────────────────────────────────────────────┘\n", + " ↓\n", + "┌──────────────────────┐ ┌───────────────────────────────────────────────────┐\n", + "│ Save Consolidated │ --> │ Store all processed metadata in a structured CSV │\n", + "│ Metadata to CSV │ │ format for final analysis. │\n", + "└──────────────────────┘ └───────────────────────────────────────────────────┘\n" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "efc9939d", + "metadata": {}, + "outputs": [], + "source": [ + "import os\n", + "import re\n", + "import json\n", + "import csv\n", + "import struct\n", + "import requests\n", + "from pathlib import Path\n", + "import pandas as pd\n", + "from collections import Counter\n", + "from concurrent.futures import ProcessPoolExecutor\n", + "from pathlib import Path\n", + "import matplotlib.pyplot as plt\n", + "from matplotlib.font_manager import FontProperties\n", + "from matplotlib import font_manager\n", + "import pandas as pd\n", + "from collections import Counter\n", + "from concurrent.futures import ProcessPoolExecutor\n", + "\n", + "# Define the current directory and important file paths\n", + "current_dir = Path.cwd()\n", + "\n", + "# Define frequently used directories\n", + "\n", + "data_dir = current_dir.parent / 'data/csv/adapters.csv'\n", + "fonts_dir = current_dir.parent / 'misc/assets/fonts'\n", + "plots_dir = current_dir.parent / 'results/plots'\n", + "raw_data_dir = current_dir.parent / 'data/adapter_metadata/lora' ### location of the LoRA metadata (JSON)\n", + "temp_dir = current_dir.parent / 'data/raw/adapters_safetensors'\n", + "misc_dir = current_dir.parent / 'misc'\n", + "\n", + "# File paths\n", + "adapters_csv = current_dir.parent / 'data/csv/adapters.csv'\n", + "output_json_dir = raw_data_dir\n", + "api_keys_file = misc_dir / 'credentials/civit.txt'\n", + "\n", + "# Ensure directories exist\n", + "os.makedirs(output_json_dir, exist_ok=True)\n", + "os.makedirs(temp_dir, exist_ok=True)\n", + "\n", + "\n", + "# Load fonts into Matplotlib\n", + "for font_path in font_paths:\n", + " font_manager.fontManager.addfont(font_path)\n", + "\n", + "# Set default font family for plots\n", + "plt.rcParams['font.family'] = ['Noto Sans JP', 'Noto Sans SC', 'sans-serif']\n", + "\n", + "print('Paths and fonts initialized successfully.')\n", + "\n", + "print('Paths initialized successfully.')" + ] + }, + { + "cell_type": "markdown", + "id": "87a58593", + "metadata": {}, + "source": [ + "## Step 2: Download LoRA and extract *.safetensors metadata\n", + "This script downloads LoRA adapters from the filtered Civiverse-Models dataset and extracts the metadata found within the *.safetensors' data structure" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "abd3a0bc", + "metadata": {}, + "outputs": [], + "source": [ + "import os\n", + "import sys\n", + "import csv\n", + "import json\n", + "import struct\n", + "import time\n", + "import requests\n", + "import signal\n", + "import contextlib\n", + "from pathlib import Path\n", + "import re\n", + "\n", + "# === Paste your API keys here ===\n", + "API_KEYS = [\n", + " \"REDACTED-CIVITAI-API-KEY-1\", #DISCORD\n", + " \"REDACTED-CIVITAI-API-KEY-2\", #ASDD 1\n", + " \"REDACTED-CIVITAI-API-KEY-3\", #ASDD 2\n", + " \"REDACTED-CIVITAI-API-KEY-4\", #BSDD \n", + " \"REDACTED-CIVITAI-API-KEY-5\"\n", + "]\n", + "if not API_KEYS or any(not isinstance(k, str) or not k.strip() for k in API_KEYS):\n", + " raise ValueError(\"Please paste at least one valid API key into API_KEYS.\")\n", + "\n", + "# === Config (adjust paths as needed) ===\n", + "current_dir = Path.cwd()\n", + "output_json_dir = current_dir.parent / \"data/adapter_metadata/lora\" # where JSON outputs go\n", + "temp_dir = current_dir.parent / \"data/raw/adapters_safetensors\" # where downloads go\n", + "csv_path = current_dir.parent / \"data/csv/adapters_poi_false_sfw.csv\"\n", + "\n", + "os.makedirs(output_json_dir, exist_ok=True)\n", + "os.makedirs(temp_dir, exist_ok=True)\n", + "\n", + "# === API key state ===\n", + "current_key_index = 0\n", + "\n", + "\n", + "\n", + "def safe_filename(name: str, max_length: int = 100) -> str:\n", + " # Replace unsafe chars\n", + " sanitized = re.sub(r'[^a-zA-Z0-9_\\-]', '_', name)\n", + " # Truncate if too long\n", + " if len(sanitized) > max_length:\n", + " sanitized = sanitized[:max_length]\n", + " return sanitized\n", + "\n", + "\n", + "def get_headers():\n", + " global current_key_index\n", + " return {\n", + " \"Accept\": \"application/json\",\n", + " \"Authorization\": f\"Bearer {API_KEYS[current_key_index].strip()}\"\n", + " }\n", + "\n", + "def rotate_api_key():\n", + " global current_key_index\n", + " if current_key_index < len(API_KEYS) - 1:\n", + " current_key_index += 1\n", + " print(f\"🔁 Rotated to API key #{current_key_index + 1}\")\n", + " else:\n", + " raise Exception(\"All API keys have been exhausted.\")\n", + "\n", + "# === Utilities ===\n", + "def save_json(data, filename):\n", + " with open(filename, 'w', encoding=\"utf-8\") as f:\n", + " json.dump(data, f, indent=4, ensure_ascii=False)\n", + "\n", + "def parse_safetensors(file_path):\n", + " # Minimal, tolerant metadata reader; returns {} on failure.\n", + " try:\n", + " with open(file_path, 'rb') as f:\n", + " file_data = f.read()\n", + " # Many safetensors use 8-byte header length; this code follows your original logic\n", + " # (4-byte) but keeps the 8-byte skip. Keep if it's working in your dataset.\n", + " metadata_size = struct.unpack(' akira\n", - " 3mma Watson v2 -> Watson\n", - " 1rene LORA -> irene\n", - " L3vi Ackerman -> Levi Ackerman\n", - "\n", - "📈 Statistics:\n", - " Total rows: 50861\n", - " Non-empty names: 50858\n", - " Empty names: 3\n", - "\n", - "🎯 Sample spaCy NER results:\n", - " 1. IU\n", - " 2. Super Pose Book\n", - " 3. Liyuu\n", - " 4. Irene\n", - " 5. AESPA Karina\n", - " 6. Saika Kawakita\n", - " 7. Liu Yifei\n", - " 8. HashimotoKanna\n", - " 9. Emma Watson\n", - " 10. Gal Gadot\n", - "\n", - "✅ Cleaned 50861 names using spaCy NER\n", - "💾 Saved to /home/lauhp/000_PHD/000_010_PUBLICATION/CODE/pm-paper/data/CSV/model_adapter/real_person_adapter_step_01_NER.csv\n" - ] - } - ], + "outputs": [], "source": [ "import pandas as pd\n", "import re\n", @@ -1011,13 +950,928 @@ "# Test mode (first 100 rows)\n", "# annotate_dataset(model_type='mistral', test_mode=True, test_size=100)\n" ] + }, + { + "cell_type": "markdown", + "id": "6431d347-d80c-4e8b-83a7-531e5df95a72", + "metadata": {}, + "source": [ + "## EuroLLM-9B-Instruct" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "e8203abc-e7c3-4cb6-aaeb-fdc6933981fc", + "metadata": {}, + "outputs": [], + "source": [ + "import pandas as pd\n", + "import json\n", + "import time\n", + "import re\n", + "from pathlib import Path\n", + "from tqdm import tqdm\n", + "import torch\n", + "from transformers import AutoModelForCausalLM, AutoTokenizer, BitsAndBytesConfig\n", + "import signal\n", + "from contextlib import contextmanager\n", + "\n", + "current_dir = Path.cwd()\n", + "input_file = current_dir.parent / \"data/CSV/model_adapter/real_person_adapter_step_02_NER.csv\"\n", + "professions_file = current_dir.parent / \"misc/lists/professions.csv\"\n", + "professions_mapped_file = current_dir.parent / \"misc/lists/professions_mapped.csv\"\n", + "# === PROCESS DATA ===\n", + "\n", + "\n", + "# === CONFIGURATION ===\n", + "TEST_MODE = False\n", + "TEST_SIZE = 100\n", + "MAX_ROWS = 50862\n", + "SAVE_INTERVAL = 10\n", + "\n", + "\n", + "index_file = current_dir.parent / \"misc/query_indicies/eurollm_local_query_index.txt\"\n", + "output_file = current_dir.parent / f\"data/CSV/eurollm_local_annotated_POI{'_test' if TEST_MODE else ''}.csv\"\n", + "\n", + "# Model settings\n", + "MODEL_NAME = \"utter-project/EuroLLM-9B-Instruct\"\n", + "#MODEL_NAME = \"Qwen/Qwen2.5-32B-Instruct\"\n", + "#MODEL_NAME = \"Qwen/Qwen2.5-14B-Instruct\"\n", + "#MODEL_NAME = \"Qwen/Qwen3-235B-A22B-Instruct-2507-FP8\"\n", + "#MODEL_NAME = \"mistralai/Mistral-7B-Instruct-v0.3\"\n", + "CACHE_DIR = current_dir.parent / \"data/models\"\n", + "CACHE_DIR.mkdir(parents=True, exist_ok=True)\n", + "\n", + "# Define the SPECIFIC profession categories\n", + "PROFESSION_CATEGORIES = [\n", + " \"actor\",\n", + " \"adult performer\",\n", + " \"singer/musician\",\n", + " \"model\",\n", + " \"online personality\",\n", + " \"public figure\",\n", + " \"voice actor/ASMR\",\n", + " \"sports professional\",\n", + " \"tv personality\"\n", + "]\n", + "\n", + "# === LOAD MODEL ===\n", + "print(f\"Loading model: {MODEL_NAME}\")\n", + "print(f\"Cache directory: {CACHE_DIR}\")\n", + "print(f\"This may take a while on first run...\\n\")\n", + "\n", + "# Check GPU availability\n", + "device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n", + "print(f\"Device: {device}\")\n", + "\n", + "if device == \"cpu\":\n", + " print(\"⚠️ WARNING: No GPU detected! Inference will be VERY slow.\")\n", + " print(\" Consider using a GPU or reducing model size.\")\n", + "\n", + "# Get HF token from credentials file\n", + "import os\n", + "credentials_dir = current_dir.parent / \"misc/credentials\"\n", + "hf_token_file = credentials_dir / \"hf_token.txt\"\n", + "\n", + "HF_TOKEN = None\n", + "if hf_token_file.exists():\n", + " HF_TOKEN = hf_token_file.read_text().strip()\n", + " print(\"✅ HF token loaded from credentials file\")\n", + "else:\n", + " print(\"⚠️ HF token file not found at:\", hf_token_file)\n", + " print(\" The script will try to use cached credentials from 'huggingface-cli login'\")\n", + " print(\" Or create the file: misc/credentials/hf_token.txt with your token\")\n", + " HF_TOKEN = None # Will use cached token if available\n", + "\n", + "# Load tokenizer\n", + "print(\"Loading tokenizer...\")\n", + "try:\n", + " tokenizer = AutoTokenizer.from_pretrained(\n", + " MODEL_NAME,\n", + " cache_dir=str(CACHE_DIR),\n", + " use_fast=True,\n", + " token=HF_TOKEN\n", + " )\n", + "except Exception as e:\n", + " print(f\"Failed with use_fast=True, trying use_fast=False...\")\n", + " tokenizer = AutoTokenizer.from_pretrained(\n", + " MODEL_NAME,\n", + " cache_dir=str(CACHE_DIR),\n", + " use_fast=False,\n", + " token=HF_TOKEN\n", + " )\n", + "\n", + "# Ensure pad token is set\n", + "if tokenizer.pad_token is None:\n", + " tokenizer.pad_token = tokenizer.eos_token\n", + "\n", + "print(\"✅ Tokenizer loaded\")\n", + "\n", + "# Configure 8-bit quantization for A100\n", + "print(\"Configuring 8-bit quantization...\")\n", + "quantization_config = BitsAndBytesConfig(\n", + " load_in_8bit=True,\n", + " llm_int8_threshold=6.0,\n", + " llm_int8_has_fp16_weight=False\n", + ")\n", + "\n", + "# Load model with 8-bit quantization\n", + "print(\"Loading model with 8-bit quantization (this may take several minutes)...\")\n", + "model = AutoModelForCausalLM.from_pretrained(\n", + " MODEL_NAME,\n", + " cache_dir=str(CACHE_DIR),\n", + " quantization_config=quantization_config,\n", + " device_map=\"auto\",\n", + " trust_remote_code=False,\n", + " token=HF_TOKEN\n", + ")\n", + "model.eval()\n", + "print(\"✅ Model loaded with 8-bit quantization\")\n", + "\n", + "# Check VRAM usage\n", + "if torch.cuda.is_available():\n", + " vram_gb = torch.cuda.max_memory_allocated() / 1024**3\n", + " print(f\"VRAM used: {vram_gb:.2f} GB\\n\")\n", + "\n", + "# === LOAD DATA ===\n", + "if output_file.exists():\n", + " print(\"Loading annotated CSV...\")\n", + " df = pd.read_csv(output_file)\n", + "else:\n", + " print(\"Loading raw input CSV...\")\n", + " df = pd.read_csv(input_file)\n", + "\n", + "\n", + "# Try to load profession mapping files\n", + "try:\n", + " professions_df = pd.read_csv(professions_file)\n", + " print(f\"✅ Loaded professions.csv\")\n", + "except:\n", + " print(\"⚠️ Warning: professions.csv not found\")\n", + "\n", + "try:\n", + " prof_mapped_df = pd.read_csv(professions_mapped_file)\n", + " print(f\"✅ Loaded profession mapping with {len(prof_mapped_df)} categories\")\n", + "except:\n", + " print(\"⚠️ Warning: professions_mapped.csv not found, using default categories\")\n", + "\n", + "profession_str = \", \".join(PROFESSION_CATEGORIES)\n", + "\n", + "print(f\"Loaded {len(df)} rows\")\n", + "print(f\"\\nProfession categories ({len(PROFESSION_CATEGORIES)}):\")\n", + "for cat in PROFESSION_CATEGORIES:\n", + " print(f\" - {cat}\")\n", + "\n", + "if TEST_MODE:\n", + " print(f\"\\nRunning in TEST MODE with {TEST_SIZE} samples\")\n", + " df = df.head(TEST_SIZE).copy()\n", + "elif MAX_ROWS:\n", + " df = df.head(MAX_ROWS).copy()\n", + "\n", + "# === CREATE PROMPTS (OPTIMIZED FOR CLEAN OUTPUTS) ===\n", + "def create_prompt(row):\n", + " \"\"\"Create prompt for EuroLLM annotation with strict formatting requirements.\"\"\"\n", + " name = row['real_name'] if pd.notna(row.get('real_name')) else row.get('name', '')\n", + " \n", + " # Gather hints\n", + " hints = []\n", + " if pd.notna(row.get('likely_profession')):\n", + " hints.append(str(row['likely_profession']))\n", + " if pd.notna(row.get('likely_nationality')):\n", + " hints.append(str(row['likely_nationality']))\n", + " if pd.notna(row.get('likely_country')):\n", + " hints.append(str(row['likely_country']))\n", + " \n", + " # Add tags if we don't have enough hints\n", + " if len(hints) < 3:\n", + " for i in range(1, 8):\n", + " tag_col = f'tag_{i}'\n", + " if tag_col in row and pd.notna(row[tag_col]):\n", + " tag_val = str(row[tag_col])\n", + " if tag_val not in hints:\n", + " hints.append(tag_val)\n", + " if len(hints) >= 5:\n", + " break\n", + " \n", + " hint_text = \", \".join(hints[:5]) if hints else \"none\"\n", + " \n", + " return f\"\"\"Extract information about '{name}'. \n", + "Context hints (DO NOT copy these as professions): {hint_text}\n", + "\n", + "Respond with EXACTLY 5 numbered lines. Each line must contain ONLY the value, no labels or extra text.\n", + "\n", + "FORMAT REQUIREMENTS:\n", + "1. Full legal name in Western order (first last). VALUE ONLY.\n", + "2. Stage names/aliases, comma-separated. If none, write \"None\". VALUE ONLY.\n", + "3. Gender: MUST be exactly one word: Male, Female, Other, or Unknown. VALUE ONLY.\n", + "4. Professions: Choose up to 3 from this list ONLY: actor, adult performer, singer/musician, model, online personality, public figure, voice actor/ASMR, sports professional, tv personality. Comma-separated. VALUE ONLY.\n", + "5. Primary country: Country name only (e.g., \"China\", \"United States\", \"Colombia\"). VALUE ONLY.\n", + "\n", + "CRITICAL RULES FOR PROFESSIONS (Line 4):\n", + "- ONLY use the exact profession categories listed above\n", + "- DO NOT use descriptive words like \"sexy\", \"photorealistic\", \"celebrity\"\n", + "- DO NOT copy the hint words as professions\n", + "- If uncertain about profession, write \"Unknown\"\n", + "- Valid professions are ONLY: actor, adult performer, singer/musician, model, online personality, public figure, voice actor/ASMR, sports professional, tv personality\n", + "- Actress = actor, streamer = online personality, YouTuber = online personality\n", + "\n", + "OTHER RULES:\n", + "- Use \"Unknown\" when uncertain or for fictional characters\n", + "- NO explanatory text, NO labels like \"Gender:\", NO prefixes\n", + "- For multi-role people, list up to 3 categories by relevance\"\"\"\n", + "\n", + "# Create prompts\n", + "print(\"\\nCreating prompts...\")\n", + "df['prompt'] = df.apply(create_prompt, axis=1)\n", + "print(\"✅ Prompts created\")\n", + "\n", + "@contextmanager\n", + "def timeout(duration):\n", + " def handler(signum, frame):\n", + " raise TimeoutError(\"Operation timed out\")\n", + " \n", + " # Set the signal handler and alarm\n", + " signal.signal(signal.SIGALRM, handler)\n", + " signal.alarm(duration)\n", + " try:\n", + " yield\n", + " finally:\n", + " signal.alarm(0) # Disable the alarm\n", + "\n", + "\n", + "def query_eurollm_local(prompt: str) -> str:\n", + " \"\"\"Query EuroLLM locally via transformers with very low temperature.\"\"\"\n", + " try:\n", + " # Format as chat message for EuroLLM with strict system prompt\n", + " messages = [\n", + " {\"role\": \"system\", \"content\": \"You are a data extraction assistant. Respond with exactly 5 numbered lines containing ONLY values. No labels, no explanations, no prefixes. Follow the format precisely.\"},\n", + " {\"role\": \"user\", \"content\": prompt}\n", + " ]\n", + " \n", + " # Tokenize\n", + " if hasattr(tokenizer, 'apply_chat_template') and tokenizer.chat_template is not None:\n", + " text = tokenizer.apply_chat_template(\n", + " messages,\n", + " tokenize=False,\n", + " add_generation_prompt=True\n", + " )\n", + " else:\n", + " # Fallback for models without chat template\n", + " text = f\"[INST] {prompt} [/INST]\"\n", + " \n", + " inputs = tokenizer([text], return_tensors=\"pt\", padding=True).to(device)\n", + " \n", + " # Generate with timeout and very low temperature\n", + " with timeout(60):\n", + " with torch.no_grad():\n", + " outputs = model.generate(\n", + " **inputs,\n", + " max_new_tokens=100,\n", + " temperature=0.01, # Very low temperature for more deterministic outputs\n", + " do_sample=True, # Must be True when temperature is set\n", + " pad_token_id=tokenizer.eos_token_id\n", + " )\n", + " \n", + " # Decode\n", + " generated_ids = outputs[0][inputs['input_ids'].shape[1]:]\n", + " response = tokenizer.decode(generated_ids, skip_special_tokens=True)\n", + " \n", + " return response.strip()\n", + " \n", + " except TimeoutError:\n", + " print(f\"[ERROR] Generation timed out after 60 seconds\")\n", + " return None\n", + " except Exception as e:\n", + " print(f\"Generation error: {e}\")\n", + " import traceback\n", + " traceback.print_exc()\n", + " return None\n", + "\n", + " \n", + "# === PARSE RESPONSE WITH CLEANING ===\n", + "def parse_response(response):\n", + " \"\"\"Parse EuroLLM response into structured fields with cleaning.\"\"\"\n", + " if not response:\n", + " return {\n", + " 'full_name': 'Unknown',\n", + " 'aliases': 'Unknown',\n", + " 'gender': 'Unknown',\n", + " 'profession_llm': 'Unknown',\n", + " 'country': 'Unknown'\n", + " }\n", + " \n", + " # Valid profession categories\n", + " VALID_PROFESSIONS = {\n", + " \"actor\", \"adult performer\", \"singer/musician\", \"model\", \n", + " \"online personality\", \"public figure\", \"voice actor/asmr\", \n", + " \"sports professional\", \"tv personality\"\n", + " }\n", + " \n", + " # Split into lines and clean\n", + " lines = [line.strip() for line in response.split('\\n') if line.strip()]\n", + " \n", + " # Initialize with Unknown values\n", + " fields = {\n", + " 'full_name': 'Unknown',\n", + " 'aliases': 'Unknown',\n", + " 'gender': 'Unknown',\n", + " 'profession_llm': 'Unknown',\n", + " 'country': 'Unknown'\n", + " }\n", + " \n", + " # Extract information from each numbered line\n", + " for line in lines:\n", + " if line.startswith('1.'):\n", + " fields['full_name'] = line[2:].strip()\n", + " elif line.startswith('2.'):\n", + " fields['aliases'] = line[2:].strip()\n", + " elif line.startswith('3.'):\n", + " # Clean gender field - remove any labels\n", + " gender_raw = line[2:].strip()\n", + " # Remove common prefixes\n", + " gender_raw = re.sub(r'^(Gender:|gender:)\\s*', '', gender_raw, flags=re.IGNORECASE)\n", + " # Extract just the gender word\n", + " gender_match = re.search(r'\\b(Male|Female|Other|Unknown)\\b', gender_raw, re.IGNORECASE)\n", + " fields['gender'] = gender_match.group(1).capitalize() if gender_match else gender_raw\n", + " elif line.startswith('4.'):\n", + " # Clean and validate profession field\n", + " profession_raw = line[2:].strip()\n", + " \n", + " # Split by comma and validate each profession\n", + " professions = [p.strip().lower() for p in profession_raw.split(',')]\n", + " valid_profs = []\n", + " \n", + " for prof in professions:\n", + " # Check if it's a valid profession\n", + " if prof in VALID_PROFESSIONS:\n", + " valid_profs.append(prof)\n", + " # Check for common invalid entries\n", + " elif prof in ['unknown', '']:\n", + " continue\n", + " # Reject descriptive words that aren't professions\n", + " elif prof in ['sexy', 'photorealistic', 'celebrity', 'famous', 'popular', \n", + " 'beautiful', 'attractive', 'hot', 'gorgeous']:\n", + " continue\n", + " # If it looks like it might be close to a valid profession, keep it\n", + " elif any(valid in prof for valid in VALID_PROFESSIONS):\n", + " # Try to extract the valid part\n", + " for valid in VALID_PROFESSIONS:\n", + " if valid in prof:\n", + " valid_profs.append(valid)\n", + " break\n", + " \n", + " # Set the cleaned professions or Unknown if none are valid\n", + " if valid_profs:\n", + " fields['profession_llm'] = ', '.join(valid_profs)\n", + " else:\n", + " fields['profession_llm'] = 'Unknown'\n", + " \n", + " elif line.startswith('5.'):\n", + " # Clean country field - remove any labels\n", + " country_raw = line[2:].strip()\n", + " # Remove common prefixes like \"Primary country:\", \"Country:\", etc.\n", + " country_raw = re.sub(r'^(Primary\\s+)?(associated\\s+)?country:\\s*', '', country_raw, flags=re.IGNORECASE)\n", + " fields['country'] = country_raw\n", + " \n", + " return fields\n", + "\n", + "# === PROCESS DATA ===\n", + "index_file.parent.mkdir(parents=True, exist_ok=True)\n", + "\n", + "# Load index\n", + "current_index = 0\n", + "if index_file.exists():\n", + " try:\n", + " current_index = int(index_file.read_text().strip())\n", + " except:\n", + " current_index = 0\n", + "\n", + "print(f\"Resuming from index {current_index}\")\n", + "\n", + "start_time = time.time()\n", + "\n", + "for i in tqdm(range(current_index, len(df)), desc=\"EuroLLM Local\"):\n", + "\n", + " prompt = df.at[i, \"prompt\"]\n", + "\n", + " # -------- MODEL QUERY WITH RETRIES --------\n", + " response = None\n", + " for attempt in range(3):\n", + " response = query_eurollm_local(prompt)\n", + " \n", + " # DEBUG: Print first few responses to see what's happening\n", + " if i < 5:\n", + " print(f\"\\n=== DEBUG Row {i}, Attempt {attempt+1} ===\")\n", + " print(f\"Response length: {len(response) if response else 0}\")\n", + " print(f\"Response: {response[:500] if response else 'None'}\")\n", + " print(\"=\" * 50)\n", + " \n", + " # Valid response?\n", + " if response and len(response.strip()) > 10:\n", + " break\n", + " \n", + " print(f\"⚠️ Row {i}: Empty or invalid response, retry {attempt+1}/3\")\n", + " time.sleep(0.5)\n", + "\n", + " # If still invalid → DO NOT overwrite previous data\n", + " if not response or len(response.strip()) <= 10:\n", + " print(f\"❌ Row {i}: failed after retries, not writing, not advancing index\")\n", + " continue\n", + "\n", + " parsed = parse_response(response)\n", + "\n", + " # DEBUG: Print first few parsed results\n", + " if i < 5:\n", + " print(f\"\\n=== PARSED Row {i} ===\")\n", + " for key, value in parsed.items():\n", + " print(f\" {key}: {value}\")\n", + " print(\"=\" * 50)\n", + "\n", + " # Additional safety: skip rows that parsed as all 'Unknown'\n", + " if all(v == \"Unknown\" for v in parsed.values()):\n", + " print(f\"❌ Row {i}: parsed as all Unknown (likely model crash); skipping.\")\n", + " continue\n", + "\n", + " # -------- WRITE PARSED FIELDS SAFELY --------\n", + " for key, value in parsed.items():\n", + " df.at[i, key] = value\n", + "\n", + " # Advance progress ONLY after successful write\n", + " current_index = i + 1\n", + "\n", + " # -------- GPU MEMORY CLEANUP --------\n", + " if torch.cuda.is_available():\n", + " torch.cuda.empty_cache()\n", + " torch.cuda.synchronize()\n", + "\n", + " # -------- SAVE LIKE YOUR DEEPSEEK VERSION --------\n", + " if (i + 1) % SAVE_INTERVAL == 0 or (i + 1) == len(df):\n", + " df.to_csv(output_file, index=False)\n", + " with open(index_file, \"w\") as f:\n", + " f.write(str(current_index))\n", + " print(f\"💾 Progress saved after row {i+1}\")\n", + "\n", + "# Final save\n", + "df.to_csv(output_file, index=False)\n", + "index_file.write_text(str(current_index))\n", + "print(\"✅ Finished full dataset.\")" + ] + }, + { + "cell_type": "markdown", + "id": "472e5ac2-ec04-4bfa-8a67-116277238c15", + "metadata": {}, + "source": [ + "## Mistral 24b instruct" + ] + }, + { + "cell_type": "code", + "execution_count": 1, + "id": "a55a5e30-83f3-4f7c-a537-b1216d4e8a07", + "metadata": { + "execution": { + "iopub.execute_input": "2025-12-08T23:57:35.685431Z", + "iopub.status.busy": "2025-12-08T23:57:35.685314Z", + "iopub.status.idle": "2025-12-08T23:59:48.656498Z", + "shell.execute_reply": "2025-12-08T23:59:48.655927Z", + "shell.execute_reply.started": "2025-12-08T23:57:35.685419Z" + } + }, + "outputs": [ + { + "name": "stderr", + "output_type": "stream", + "text": [ + "/shares/weddigen.ki.uzh/laura_wagner/phase_01/pm-paper/.venv/lib/python3.11/site-packages/tqdm/auto.py:21: TqdmWarning: IProgress not found. Please update jupyter and ipywidgets. See https://ipywidgets.readthedocs.io/en/stable/user_install.html\n", + " from .autonotebook import tqdm as notebook_tqdm\n" + ] + }, + { + "name": "stdout", + "output_type": "stream", + "text": [ + "Loading model: mistralai/Mistral-Small-Instruct-2409\n", + "Cache directory: /shares/weddigen.ki.uzh/laura_wagner/phase_01/pm-paper/data/models\n", + "This may take a while on first run (~65GB download)...\n", + "\n", + "Device: cuda\n", + "Loading tokenizer...\n", + "✅ Tokenizer loaded\n" + ] + }, + { + "ename": "NameError", + "evalue": "name 'BitsAndBytesConfig' is not defined", + "output_type": "error", + "traceback": [ + "\u001b[31m---------------------------------------------------------------------------\u001b[39m", + "\u001b[31mNameError\u001b[39m Traceback (most recent call last)", + "\u001b[36mCell\u001b[39m\u001b[36m \u001b[39m\u001b[32mIn[1]\u001b[39m\u001b[32m, line 82\u001b[39m\n\u001b[32m 78\u001b[39m tokenizer.pad_token = tokenizer.eos_token\n\u001b[32m 80\u001b[39m \u001b[38;5;28mprint\u001b[39m(\u001b[33m\"\u001b[39m\u001b[33m✅ Tokenizer loaded\u001b[39m\u001b[33m\"\u001b[39m)\n\u001b[32m---> \u001b[39m\u001b[32m82\u001b[39m quantization_config = \u001b[43mBitsAndBytesConfig\u001b[49m(\n\u001b[32m 83\u001b[39m load_in_8bit=\u001b[38;5;28;01mTrue\u001b[39;00m\n\u001b[32m 84\u001b[39m )\n\u001b[32m 87\u001b[39m \u001b[38;5;66;03m# Load model with optimizations\u001b[39;00m\n\u001b[32m 88\u001b[39m \u001b[38;5;28mprint\u001b[39m(\u001b[33m\"\u001b[39m\u001b[33mLoading model (this may take several minutes)...\u001b[39m\u001b[33m\"\u001b[39m)\n", + "\u001b[31mNameError\u001b[39m: name 'BitsAndBytesConfig' is not defined" + ] + } + ], + "source": [ + "import pandas as pd\n", + "import json\n", + "import time\n", + "import re\n", + "from pathlib import Path\n", + "from tqdm import tqdm\n", + "import torch\n", + "from transformers import AutoModelForCausalLM, AutoTokenizer\n", + "\n", + "current_dir = Path.cwd()\n", + "input_file = current_dir.parent / \"data/CSV/model_adapter/real_person_adapter_step_02_NER.csv\"\n", + "professions_file = current_dir.parent / \"misc/lists/professions.csv\"\n", + "professions_mapped_file = current_dir.parent / \"misc/lists/professions_mapped.csv\"\n", + "# === PROCESS DATA ===\n", + "\n", + "\n", + "# === CONFIGURATION ===\n", + "TEST_MODE = False\n", + "TEST_SIZE = 100\n", + "MAX_ROWS = 50862\n", + "SAVE_INTERVAL = 10\n", + "\n", + "output_file = current_dir.parent / f\"data/CSV/mistral24_local_annotated_POI{'_test' if TEST_MODE else ''}.csv\"\n", + "index_file = current_dir.parent / \"misc/query_indicies/mistral24_local_query_index.txt\"\n", + "\n", + "\n", + "# Model settings\n", + "#MODEL_NAME = \"mistralai/Mistral-Small-3.1-24B-Instruct-2503\"\n", + "MODEL_NAME = \"mistralai/Mistral-Small-Instruct-2409\"\n", + "#MODEL_NAME = \"mistralai/Mistral-7B-Instruct-v0.3\"\n", + "CACHE_DIR = current_dir.parent / \"data/models\"\n", + "CACHE_DIR.mkdir(parents=True, exist_ok=True)\n", + "\n", + "# Define the SPECIFIC profession categories\n", + "PROFESSION_CATEGORIES = [\n", + " \"actor\",\n", + " \"adult performer\",\n", + " \"singer/musician\",\n", + " \"model\",\n", + " \"online personality\",\n", + " \"public figure\",\n", + " \"voice actor/ASMR\",\n", + " \"sports professional\",\n", + " \"tv personality\"\n", + "]\n", + "\n", + "# === LOAD MODEL ===\n", + "print(f\"Loading model: {MODEL_NAME}\")\n", + "print(f\"Cache directory: {CACHE_DIR}\")\n", + "print(f\"This may take a while on first run (~65GB download)...\\n\")\n", + "\n", + "# Check GPU availability\n", + "device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n", + "print(f\"Device: {device}\")\n", + "\n", + "if device == \"cpu\":\n", + " print(\"⚠️ WARNING: No GPU detected! Inference will be VERY slow.\")\n", + " print(\" Consider using a GPU or reducing model size.\")\n", + "\n", + "# Load tokenizer\n", + "print(\"Loading tokenizer...\")\n", + "try:\n", + " tokenizer = AutoTokenizer.from_pretrained(\n", + " MODEL_NAME,\n", + " cache_dir=str(CACHE_DIR),\n", + " use_fast=True\n", + " )\n", + "except Exception as e:\n", + " print(f\"Failed with use_fast=True, trying use_fast=False...\")\n", + " tokenizer = AutoTokenizer.from_pretrained(\n", + " MODEL_NAME,\n", + " cache_dir=str(CACHE_DIR),\n", + " use_fast=False\n", + " )\n", + "\n", + "# Ensure pad token is set\n", + "if tokenizer.pad_token is None:\n", + " tokenizer.pad_token = tokenizer.eos_token\n", + "\n", + "print(\"✅ Tokenizer loaded\")\n", + "\n", + "quantization_config = BitsAndBytesConfig(\n", + " load_in_8bit=True\n", + ")\n", + "\n", + "\n", + "# Load model with optimizations\n", + "print(\"Loading model (this may take several minutes)...\")\n", + "model = AutoModelForCausalLM.from_pretrained(\n", + " MODEL_NAME,\n", + " cache_dir=str(CACHE_DIR),\n", + " torch_dtype=torch.bfloat16,\n", + " quantization_config=quantization_config,\n", + " device_map=\"auto\",\n", + " trust_remote_code=False\n", + ")\n", + "model.eval()\n", + "print(\"✅ Model loaded\")\n", + "\n", + "# Check VRAM usage\n", + "if torch.cuda.is_available():\n", + " vram_gb = torch.cuda.max_memory_allocated() / 1024**3\n", + " print(f\"VRAM used: {vram_gb:.2f} GB\\n\")\n", + "\n", + "# === LOAD DATA ===\n", + "print(\"Loading raw input CSV...\")\n", + "df = pd.read_csv(input_file) # ALWAYS load the full input\n", + "print(f\"Loaded {len(df)} rows from input file\")\n", + "\n", + "# If we have previous annotations, merge them\n", + "if output_file.exists():\n", + " print(\"Found existing annotations, merging...\")\n", + " existing_df = pd.read_csv(output_file)\n", + " print(f\"Existing annotations has {len(existing_df)} rows\")\n", + " \n", + " # Update df with existing annotations\n", + " # Only update the columns that were annotated\n", + " annotation_cols = ['full_name', 'aliases', 'gender', 'profession_llm', 'country']\n", + " for col in annotation_cols:\n", + " if col in existing_df.columns:\n", + " df[col] = existing_df[col][:len(df)] # Make sure we don't exceed df length\n", + " \n", + " print(f\"Merged annotations, continuing with {len(df)} total rows\")\n", + "\n", + "\n", + "# Try to load profession mapping files\n", + "try:\n", + " professions_df = pd.read_csv(professions_file)\n", + " print(f\"✅ Loaded professions.csv\")\n", + "except:\n", + " print(\"⚠️ Warning: professions.csv not found\")\n", + "\n", + "try:\n", + " prof_mapped_df = pd.read_csv(professions_mapped_file)\n", + " print(f\"✅ Loaded profession mapping with {len(prof_mapped_df)} categories\")\n", + "except:\n", + " print(\"⚠️ Warning: professions_mapped.csv not found, using default categories\")\n", + "\n", + "profession_str = \", \".join(PROFESSION_CATEGORIES)\n", + "\n", + "print(f\"Loaded {len(df)} rows\")\n", + "print(f\"\\nProfession categories ({len(PROFESSION_CATEGORIES)}):\")\n", + "for cat in PROFESSION_CATEGORIES:\n", + " print(f\" - {cat}\")\n", + "\n", + "if TEST_MODE:\n", + " print(f\"\\nRunning in TEST MODE with {TEST_SIZE} samples\")\n", + " df = df.head(TEST_SIZE).copy()\n", + "elif MAX_ROWS:\n", + " df = df.head(MAX_ROWS).copy()\n", + "\n", + "# === CREATE PROMPTS (DEEPSEEK STYLE) ===\n", + "def create_prompt(row):\n", + " \"\"\"Create prompt for Mistral annotation with specific profession categories.\"\"\"\n", + " name = row['real_name'] if pd.notna(row.get('real_name')) else row.get('name', '')\n", + " \n", + " # Gather hints\n", + " hints = []\n", + " if pd.notna(row.get('likely_profession')):\n", + " hints.append(str(row['likely_profession']))\n", + " if pd.notna(row.get('likely_nationality')):\n", + " hints.append(str(row['likely_nationality']))\n", + " if pd.notna(row.get('likely_country')):\n", + " hints.append(str(row['likely_country']))\n", + " \n", + " # Add tags if we don't have enough hints\n", + " if len(hints) < 3:\n", + " for i in range(1, 8):\n", + " tag_col = f'tag_{i}'\n", + " if tag_col in row and pd.notna(row[tag_col]):\n", + " tag_val = str(row[tag_col])\n", + " if tag_val not in hints:\n", + " hints.append(tag_val)\n", + " if len(hints) >= 5:\n", + " break\n", + " \n", + " hint_text = \", \".join(hints[:5]) if hints else \"none\"\n", + " \n", + " return f\"\"\"Given '{name}' ({hint_text}), provide:\n", + "1. Full legal name (Western order if non-latin script)\n", + "2. Any stage names/aliases (comma separated)\n", + "3. Gender (Male/Female/Other/Unknown)\n", + "4. Top 3 most likely professions from ONLY these categories:\n", + " - actor\n", + " - adult performer\n", + " - singer/musician\n", + " - model\n", + " - online personality (includes streamers, cosplayers, influencers)\n", + " - public figure (includes politicians, activists, journalists, authors)\n", + " - voice actor/ASMR\n", + " - sports professional\n", + " - tv personality (includes hosts, presenters, reality TV)\n", + "\n", + "5. Primary country associated\n", + "\n", + "IMPORTANT:\n", + "- Choose professions ONLY from the 9 categories above\n", + "- Provide up to 3 professions, comma-separated, ordered by relevance\n", + "- Be SPECIFIC: choose the most accurate category for each role\n", + "- \"online personality\" includes: streamers, cosplayers, YouTubers, influencers, content creators\n", + "- Use 'Unknown' when uncertain or for fictional characters/places\n", + "- For multi-role people, list all relevant categories (e.g., \"actor, singer/musician, online personality\")\n", + "- For country respond with one word only, for example China or Columbia\n", + "- actress = actor\n", + "\n", + "Respond with exactly 5 numbered lines.\"\"\"\n", + "\n", + "# Create prompts\n", + "print(\"\\nCreating prompts...\")\n", + "df['prompt'] = df.apply(create_prompt, axis=1)\n", + "print(\"✅ Prompts created\")\n", + "\n", + "# === QUERY MISTRAL LOCAL ===\n", + "def query_mistral_local(prompt: str) -> str:\n", + " \"\"\"Query Mistral locally via transformers.\"\"\"\n", + " try:\n", + " # Format as chat message for Mistral\n", + " messages = [\n", + " {\"role\": \"system\", \"content\": \"You are an assistant that extracts key data on a person based on the name. Respond with exactly 5 numbered lines. For professions, choose ONLY from these categories: actor, adult performer, singer/musician, model, online personality, public figure, voice actor/ASMR, sports professional, tv personality.\"},\n", + " {\"role\": \"user\", \"content\": prompt}\n", + " ]\n", + " \n", + " # Tokenize\n", + " if hasattr(tokenizer, 'apply_chat_template'):\n", + " text = tokenizer.apply_chat_template(\n", + " messages,\n", + " tokenize=False,\n", + " add_generation_prompt=True\n", + " )\n", + " else:\n", + " # Fallback for older tokenizers\n", + " text = f\"[INST] {prompt} [/INST]\"\n", + " \n", + " inputs = tokenizer([text], return_tensors=\"pt\", padding=True).to(device)\n", + " \n", + " # Generate\n", + " with torch.no_grad():\n", + " outputs = model.generate(\n", + " **inputs,\n", + " max_new_tokens=512,\n", + " temperature=0.05,\n", + " do_sample=True,\n", + " top_p=0.8,\n", + " pad_token_id=tokenizer.pad_token_id if tokenizer.pad_token_id else tokenizer.eos_token_id\n", + " )\n", + " \n", + " # Decode\n", + " generated_ids = outputs[0][inputs['input_ids'].shape[1]:]\n", + " response = tokenizer.decode(generated_ids, skip_special_tokens=True)\n", + " \n", + " return response.strip()\n", + " \n", + " except Exception as e:\n", + " print(f\"Generation error: {e}\")\n", + " return None\n", + "\n", + "# === PARSE RESPONSE (DEEPSEEK STYLE) ===\n", + "def parse_response(response):\n", + " \"\"\"Parse Mistral response into structured fields.\"\"\"\n", + " if not response:\n", + " return {\n", + " 'full_name': 'Unknown',\n", + " 'aliases': 'Unknown',\n", + " 'gender': 'Unknown',\n", + " 'profession_llm': 'Unknown',\n", + " 'country': 'Unknown'\n", + " }\n", + " \n", + " # Split into lines and clean\n", + " lines = [line.strip() for line in response.split('\\n') if line.strip()]\n", + " \n", + " # Initialize with Unknown values\n", + " fields = {\n", + " 'full_name': 'Unknown',\n", + " 'aliases': 'Unknown',\n", + " 'gender': 'Unknown',\n", + " 'profession_llm': 'Unknown',\n", + " 'country': 'Unknown'\n", + " }\n", + " \n", + " # Extract information from each numbered line\n", + " for line in lines:\n", + " if line.startswith('1.'):\n", + " fields['full_name'] = line[2:].strip()\n", + " elif line.startswith('2.'):\n", + " fields['aliases'] = line[2:].strip()\n", + " elif line.startswith('3.'):\n", + " fields['gender'] = line[2:].strip()\n", + " elif line.startswith('4.'):\n", + " fields['profession_llm'] = line[2:].strip()\n", + " elif line.startswith('5.'):\n", + " fields['country'] = line[2:].strip()\n", + " \n", + " return fields\n", + "\n", + "# === PROCESS DATA ===\n", + "output_file = current_dir.parent / f\"data/CSV/mistral24_local_annotated_POI{'_test' if TEST_MODE else ''}.csv\"\n", + "index_file = current_dir.parent / \"misc/query_indicies/mistral24_local_query_index.txt\"\n", + "\n", + "index_file.parent.mkdir(parents=True, exist_ok=True)\n", + "\n", + "# Load index\n", + "current_index = 0\n", + "if index_file.exists():\n", + " try:\n", + " current_index = int(index_file.read_text().strip())\n", + " except:\n", + " current_index = 0\n", + "\n", + "print(f\"Resuming from index {current_index}\")\n", + "\n", + "start_time = time.time()\n", + "\n", + "for i in tqdm(range(current_index, len(df)), desc=\"Mistral Local\"):\n", + "\n", + " prompt = df.at[i, \"prompt\"]\n", + "\n", + " # -------- MODEL QUERY WITH RETRIES --------\n", + " response = None\n", + " for attempt in range(3):\n", + " response = query_mistral_local(prompt)\n", + " \n", + " # Valid response?\n", + " if response and len(response.strip()) > 10:\n", + " break\n", + " \n", + " print(f\"⚠️ Row {i}: Empty or invalid response, retry {attempt+1}/3\")\n", + " time.sleep(0.5)\n", + "\n", + " # If still invalid → DO NOT overwrite previous data\n", + " if not response or len(response.strip()) <= 10:\n", + " print(f\"❌ Row {i}: failed after retries, not writing, not advancing index\")\n", + " continue\n", + "\n", + " parsed = parse_response(response)\n", + "\n", + " # Additional safety: skip rows that parsed as all 'Unknown'\n", + " if all(v == \"Unknown\" for v in parsed.values()):\n", + " print(f\"❌ Row {i}: parsed as all Unknown (likely model crash); skipping.\")\n", + " continue\n", + "\n", + " # -------- WRITE PARSED FIELDS SAFELY --------\n", + " for key, value in parsed.items():\n", + " df.at[i, key] = value\n", + "\n", + " # Advance progress ONLY after successful write\n", + " current_index = i + 1\n", + "\n", + " # -------- GPU MEMORY CLEANUP --------\n", + " if torch.cuda.is_available():\n", + " torch.cuda.empty_cache()\n", + " torch.cuda.synchronize()\n", + "\n", + " # -------- SAVE LIKE YOUR DEEPSEEK VERSION --------\n", + " if (i + 1) % SAVE_INTERVAL == 0 or (i + 1) == len(df):\n", + " df.to_csv(output_file, index=False)\n", + " with open(index_file, \"w\") as f:\n", + " f.write(str(current_index))\n", + " print(f\"💾 Progress saved after row {i+1}\")\n", + "\n", + "# Final save\n", + "df.to_csv(output_file, index=False)\n", + "index_file.write_text(str(current_index))\n", + "print(\"✅ Finished full dataset.\")\n" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "d7212e75-0ff6-45a0-8695-c4a3d3e02818", + "metadata": {}, + "outputs": [], + "source": [ + "import transformers\n", + "print(f\"Transformers version: {transformers.__version__}\")\n", + "\n", + "# Check if Mistral3 is available\n", + "try:\n", + " from transformers import Mistral3ForCausalLM\n", + " print(\"✅ Mistral3 is available\")\n", + "except ImportError:\n", + " print(\"❌ Mistral3 not available in this transformers version\")" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "a6ab032e-246e-4c4e-9776-ff0bfbf6fd9c", + "metadata": {}, + "outputs": [], + "source": [] } ], "metadata": { "kernelspec": { - "display_name": "latm", + "display_name": "pm-paper", "language": "python", - "name": "python3" + "name": "pm-paper" }, "language_info": { "codemirror_mode": { @@ -1029,7 +1883,7 @@ "name": "python", "nbconvert_exporter": "python", "pygments_lexer": "ipython3", - "version": "3.10.15" + "version": "3.11.13" } }, "nbformat": 4, diff --git a/jupyter_notebooks/Section_2-4_Figure_9_ectract_LoRA_metadata_v2.ipynb b/jupyter_notebooks/Section_2-4_Figure_9_ectract_LoRA_metadata_v2.ipynb index 7ad7463163f1d7a7ffdc81a2de86f326c47c551d..728cf9f3ec66eadda045ba883a06a6443fa83cec 100644 --- a/jupyter_notebooks/Section_2-4_Figure_9_ectract_LoRA_metadata_v2.ipynb +++ b/jupyter_notebooks/Section_2-4_Figure_9_ectract_LoRA_metadata_v2.ipynb @@ -391,8 +391,22 @@ } ], "metadata": { + "kernelspec": { + "display_name": "Python 3 (ipykernel)", + "language": "python", + "name": "python3" + }, "language_info": { - "name": "python" + "codemirror_mode": { + "name": "ipython", + "version": 3 + }, + "file_extension": ".py", + "mimetype": "text/x-python", + "name": "python", + "nbconvert_exporter": "python", + "pygments_lexer": "ipython3", + "version": "3.13.9" } }, "nbformat": 4, diff --git a/misc/query_indicies/eurollm_local_query_index.txt b/misc/query_indicies/eurollm_local_query_index.txt new file mode 100644 index 0000000000000000000000000000000000000000..e3f1e9b791c84fce95fe992dc246e9e2286c84ed --- /dev/null +++ b/misc/query_indicies/eurollm_local_query_index.txt @@ -0,0 +1 @@ +80 \ No newline at end of file diff --git a/misc/query_indicies/mistral24_local_query_index.txt b/misc/query_indicies/mistral24_local_query_index.txt new file mode 100644 index 0000000000000000000000000000000000000000..0d2f460b7b86b185f179023ab99dbed8233d46bf --- /dev/null +++ b/misc/query_indicies/mistral24_local_query_index.txt @@ -0,0 +1 @@ +3890 \ No newline at end of file