laura.wagner commited on
Commit
4939c15
·
1 Parent(s): 5fabb43

added mistral 24b code

Browse files
jupyter_notebooks/.ipynb_checkpoints/Section_2-3-4_Figure_8_Step_1_LLM_annotation-checkpoint.ipynb ADDED
@@ -0,0 +1,1451 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "cells": [
3
+ {
4
+ "cell_type": "markdown",
5
+ "id": "23d0ae58",
6
+ "metadata": {},
7
+ "source": [
8
+ "# Deepfake Adapter Dataset - LLM Annotation Pipeline"
9
+ ]
10
+ },
11
+ {
12
+ "cell_type": "markdown",
13
+ "id": "e4407358",
14
+ "metadata": {},
15
+ "source": [
16
+ "### Unified Model Loading & Inference\n",
17
+ "Code for querying Mistral, Gemma, and Qwen models."
18
+ ]
19
+ },
20
+ {
21
+ "cell_type": "markdown",
22
+ "id": "1a1b9d0e",
23
+ "metadata": {},
24
+ "source": [
25
+ "## CLEANING & PREPROCESSING"
26
+ ]
27
+ },
28
+ {
29
+ "cell_type": "markdown",
30
+ "id": "3df42c46",
31
+ "metadata": {},
32
+ "source": [
33
+ "#### Named Entity Recognitition (NER) using SpaCy "
34
+ ]
35
+ },
36
+ {
37
+ "cell_type": "code",
38
+ "execution_count": null,
39
+ "id": "a287eef4",
40
+ "metadata": {},
41
+ "outputs": [],
42
+ "source": [
43
+ "import pandas as pd\n",
44
+ "import re\n",
45
+ "from pathlib import Path\n",
46
+ "import emoji\n",
47
+ "import spacy\n",
48
+ "\n",
49
+ "# Load spaCy model\n",
50
+ "# You may need to download it first: python -m spacy download en_core_web_sm\n",
51
+ "try:\n",
52
+ " nlp = spacy.load(\"en_core_web_sm\")\n",
53
+ " print(\"✅ spaCy model loaded: en_core_web_sm\")\n",
54
+ "except OSError:\n",
55
+ " print(\"❌ spaCy model not found. Downloading...\")\n",
56
+ " import subprocess\n",
57
+ " subprocess.run([\"python\", \"-m\", \"spacy\", \"download\", \"en_core_web_sm\"])\n",
58
+ " nlp = spacy.load(\"en_core_web_sm\")\n",
59
+ " print(\"✅ spaCy model downloaded and loaded\")\n",
60
+ "\n",
61
+ "# Set up paths\n",
62
+ "current_dir = Path.cwd()\n",
63
+ "#input_file = current_dir.parent / \"data/CSV/real_person_adapters.csv\"\n",
64
+ "input_file = current_dir.parent / \"data/CSV/model_adapter/real_person_adapter.csv\"\n",
65
+ "\n",
66
+ "# Load dataset\n",
67
+ "df = pd.read_csv(input_file)\n",
68
+ "print(f\"Loaded {len(df)} rows\")\n",
69
+ "\n",
70
+ "def translate_leetspeak(text: str) -> str:\n",
71
+ " \"\"\"\n",
72
+ " Translate common leetspeak patterns to normal letters.\n",
73
+ " Examples: 4kira -> Akira, 3mma -> Emma, 1rene -> Irene\n",
74
+ " \"\"\"\n",
75
+ " if not text:\n",
76
+ " return text\n",
77
+ " \n",
78
+ " # Common leetspeak mappings (order matters!)\n",
79
+ " leetspeak_map = {\n",
80
+ " '4': 'a',\n",
81
+ " '3': 'e', \n",
82
+ " '1': 'i',\n",
83
+ " '0': 'o',\n",
84
+ " '7': 't',\n",
85
+ " '5': 's',\n",
86
+ " '8': 'b',\n",
87
+ " '9': 'g',\n",
88
+ " '@': 'a',\n",
89
+ " '$': 's',\n",
90
+ " '!': 'i',\n",
91
+ " }\n",
92
+ " \n",
93
+ " result = text\n",
94
+ " # Apply mappings at word boundaries or start of string\n",
95
+ " for leet, normal in leetspeak_map.items():\n",
96
+ " # Replace at start of word\n",
97
+ " result = re.sub(rf'\\b{re.escape(leet)}', normal, result, flags=re.IGNORECASE)\n",
98
+ " # Replace standalone numbers that look like letters in context\n",
99
+ " result = re.sub(rf'(?<=[a-z]){re.escape(leet)}(?=[a-z])', normal, result, flags=re.IGNORECASE)\n",
100
+ " \n",
101
+ " return result\n",
102
+ "\n",
103
+ "def preprocess_for_ner(name: str) -> str:\n",
104
+ " \"\"\"\n",
105
+ " Preprocess the name before spaCy NER.\n",
106
+ " Remove noise but keep the actual name parts.\n",
107
+ " \"\"\"\n",
108
+ " if pd.isna(name):\n",
109
+ " return \"\"\n",
110
+ " \n",
111
+ " name = str(name)\n",
112
+ " \n",
113
+ " # FIRST: Translate leetspeak\n",
114
+ " name = translate_leetspeak(name)\n",
115
+ " \n",
116
+ " # Remove emoji\n",
117
+ " name = emoji.replace_emoji(name, replace=' ')\n",
118
+ " \n",
119
+ " # Remove version indicators (v1, v2, v1.0, etc.)\n",
120
+ " name = re.sub(r'\\s*[vV]\\d+(\\.\\d+)?\\s*', ' ', name)\n",
121
+ " \n",
122
+ " # Remove LoRA-related terms (case insensitive)\n",
123
+ " lora_terms = ['lora', 'loha', 'lycoris', 'controlnet', 'textual inversion', \n",
124
+ " 'embedding', 'ti', 'checkpoint', 'model', 'adapter', 'pony', 'sdxl', 'flux', 'illustrious', 'sd14', 'sd14', 'sd2', 'sd3', 'diffusion', 'stable', 'hunyuan']\n",
125
+ " for term in lora_terms:\n",
126
+ " name = re.sub(rf'\\b{term}\\b', '', name, flags=re.IGNORECASE)\n",
127
+ " \n",
128
+ " # Remove content in parentheses or brackets (often metadata)\n",
129
+ " name = re.sub(r'\\([^)]*\\)', '', name)\n",
130
+ " name = re.sub(r'\\[[^\\]]*\\]', '', name)\n",
131
+ " \n",
132
+ " # Remove special characters like 「」\n",
133
+ " name = re.sub(r'[「」『』【】〈〉《》]', '', name)\n",
134
+ " \n",
135
+ " # Handle pipe - keep first part\n",
136
+ " if '|' in name:\n",
137
+ " name = name.split('|')[0]\n",
138
+ " \n",
139
+ " # Handle forward slash - keep first part\n",
140
+ " if '/' in name:\n",
141
+ " name = name.split('/')[0]\n",
142
+ " \n",
143
+ " # Replace underscores with spaces\n",
144
+ " name = name.replace('_', ' ')\n",
145
+ " \n",
146
+ " # Remove multiple spaces\n",
147
+ " name = re.sub(r'\\s+', ' ', name)\n",
148
+ " \n",
149
+ " # Strip\n",
150
+ " name = name.strip()\n",
151
+ " \n",
152
+ " return name\n",
153
+ "\n",
154
+ "def extract_person_name(text: str) -> str:\n",
155
+ " \"\"\"\n",
156
+ " Use spaCy NER to extract person names from text.\n",
157
+ " Falls back to cleaned text if no PERSON entity found.\n",
158
+ " \"\"\"\n",
159
+ " if not text:\n",
160
+ " return \"\"\n",
161
+ " \n",
162
+ " # Run spaCy NER\n",
163
+ " doc = nlp(text)\n",
164
+ " \n",
165
+ " # Extract PERSON entities\n",
166
+ " person_entities = [ent.text for ent in doc.ents if ent.label_ == \"PERSON\"]\n",
167
+ " \n",
168
+ " if person_entities:\n",
169
+ " # Return the first (usually longest/best) person name\n",
170
+ " return person_entities[0].strip()\n",
171
+ " \n",
172
+ " # If no PERSON entity found, try to extract capitalized words (likely names)\n",
173
+ " # This helps with names spaCy might miss\n",
174
+ " words = text.split()\n",
175
+ " capitalized_words = [w for w in words if w and w[0].isupper() and len(w) > 1]\n",
176
+ " \n",
177
+ " if capitalized_words:\n",
178
+ " # Join first few capitalized words (likely the name)\n",
179
+ " return ' '.join(capitalized_words[:3]).strip()\n",
180
+ " \n",
181
+ " # Last resort: return cleaned text\n",
182
+ " return text.strip()\n",
183
+ "\n",
184
+ "def clean_name_with_spacy(name: str) -> str:\n",
185
+ " \"\"\"\n",
186
+ " Complete name cleaning pipeline with spaCy NER.\n",
187
+ " \n",
188
+ " Pipeline:\n",
189
+ " 1. Translate leetspeak (4→a, 3→e, 1→i, etc.)\n",
190
+ " 2. Remove noise (emoji, version tags, LoRA terms)\n",
191
+ " 3. Use spaCy to extract PERSON entities\n",
192
+ " 4. Fallback to capitalized words or cleaned text\n",
193
+ " \"\"\"\n",
194
+ " # Step 1 & 2: Preprocess (leetspeak + noise removal)\n",
195
+ " preprocessed = preprocess_for_ner(name)\n",
196
+ " \n",
197
+ " if not preprocessed:\n",
198
+ " return \"\"\n",
199
+ " \n",
200
+ " # Step 3: Extract person name using spaCy NER\n",
201
+ " person_name = extract_person_name(preprocessed)\n",
202
+ " \n",
203
+ " return person_name\n",
204
+ "\n",
205
+ "# Apply name cleaning with spaCy\n",
206
+ "print(\"\\n🔄 Processing names with spaCy NER...\")\n",
207
+ "df['real_name'] = df['name'].apply(clean_name_with_spacy)\n",
208
+ "\n",
209
+ "# Show examples with detailed comparison\n",
210
+ "print(\"\\n📊 Name cleaning examples (with spaCy NER):\")\n",
211
+ "print(\"=\" * 100)\n",
212
+ "print(f\"{'Original Name':<50} | {'Cleaned Name':<30}\")\n",
213
+ "print(\"=\" * 100)\n",
214
+ "\n",
215
+ "examples = df[['name', 'real_name']].head(30)\n",
216
+ "shown = 0\n",
217
+ "for idx, row in examples.iterrows():\n",
218
+ " if row['name'] != row['real_name'] and shown < 20:\n",
219
+ " print(f\"{row['name']:<50} | {row['real_name']:<30}\")\n",
220
+ " shown += 1\n",
221
+ "\n",
222
+ "print(\"=\" * 100)\n",
223
+ "\n",
224
+ "# Show specific test cases\n",
225
+ "print(\"\\n🧪 Leetspeak translation examples:\")\n",
226
+ "test_names = ['4kira LoRA', '3mma Watson v2', '1rene LORA', 'L3vi Ackerman']\n",
227
+ "for test in test_names:\n",
228
+ " result = clean_name_with_spacy(test)\n",
229
+ " print(f\" {test:<30} -> {result}\")\n",
230
+ "\n",
231
+ "# Statistics\n",
232
+ "print(f\"\\n📈 Statistics:\")\n",
233
+ "print(f\" Total rows: {len(df)}\")\n",
234
+ "print(f\" Non-empty names: {(df['real_name'] != '').sum()}\")\n",
235
+ "print(f\" Empty names: {(df['real_name'] == '').sum()}\")\n",
236
+ "\n",
237
+ "# Show some examples of what spaCy identified\n",
238
+ "print(\"\\n🎯 Sample spaCy NER results:\")\n",
239
+ "sample_names = df['real_name'].head(20).tolist()\n",
240
+ "for i, name in enumerate(sample_names[:10], 1):\n",
241
+ " if name:\n",
242
+ " print(f\" {i}. {name}\")\n",
243
+ "\n",
244
+ "print(f\"\\n✅ Cleaned {len(df)} names using spaCy NER\")\n",
245
+ "\n",
246
+ "# Save intermediate result\n",
247
+ "output_step1 = current_dir.parent / \"data/CSV/model_adapter/real_person_adapter_step_01_NER.csv\"\n",
248
+ "df.to_csv(output_step1, index=False)\n",
249
+ "print(f\"💾 Saved to {output_step1}\")\n"
250
+ ]
251
+ },
252
+ {
253
+ "cell_type": "markdown",
254
+ "id": "64687c72",
255
+ "metadata": {},
256
+ "source": [
257
+ "#### STEP 02: Nationality tag to Country hint\n",
258
+ "here tags related to nationality gets converted to the country equivalent."
259
+ ]
260
+ },
261
+ {
262
+ "cell_type": "code",
263
+ "execution_count": null,
264
+ "id": "d6eaef5b",
265
+ "metadata": {},
266
+ "outputs": [],
267
+ "source": [
268
+ "import pandas as pd\n",
269
+ "from pathlib import Path\n",
270
+ "\n",
271
+ "# Set up paths\n",
272
+ "current_dir = Path.cwd()\n",
273
+ "countries_file = current_dir.parent / \"misc/lists/countries.csv\"\n",
274
+ "professions_file = current_dir.parent / \"misc/lists/professions.csv\"\n",
275
+ "input_file = current_dir.parent / \"data/CSV/model_adapter/real_person_adapter_step_01_NER.csv\"\n",
276
+ "output_file = current_dir.parent / \"data/CSV/model_adapter/real_person_adapter_step_02_NER.csv\"\n",
277
+ "\n",
278
+ "# Load datasets\n",
279
+ "poi_df = pd.read_csv(input_file)\n",
280
+ "countries_df = pd.read_csv(countries_file)\n",
281
+ "professions_df = pd.read_csv(professions_file)\n",
282
+ "\n",
283
+ "# Define uninhabited or non-relevant territories to exclude\n",
284
+ "excluded_territories = {\n",
285
+ " 'isle of man', 'bouvet island', 'heard island and mcdonald islands',\n",
286
+ " 'french southern territories', 'south georgia and the south sandwich islands',\n",
287
+ " 'svalbard and jan mayen', 'british indian ocean territory', 'antarctica',\n",
288
+ " 'christmas island', 'cocos (keeling) islands', 'norfolk island',\n",
289
+ " 'pitcairn', 'tokelau', 'united states minor outlying islands',\n",
290
+ " 'wallis and futuna', 'western sahara'\n",
291
+ "}\n",
292
+ "\n",
293
+ "# Step 1: Combine tags into one lowercase list\n",
294
+ "def combine_tags(row):\n",
295
+ " return [str(row[f\"tag_{i}\"]).strip().lower() for i in range(1, 8) if pd.notna(row.get(f\"tag_{i}\"))]\n",
296
+ "\n",
297
+ "poi_df[\"tags\"] = poi_df.apply(combine_tags, axis=1)\n",
298
+ "\n",
299
+ "# Step 2: Build tag → (country, nationality) mapping with PRIORITIES\n",
300
+ "tag_to_country_nationality = {}\n",
301
+ "# We'll use a priority score: direct country name = 3, nationality = 2, word parts = 1\n",
302
+ "\n",
303
+ "for _, row in countries_df.iterrows():\n",
304
+ " country = str(row[\"en_short_name\"]).strip()\n",
305
+ " nationality = str(row[\"nationality\"]).strip()\n",
306
+ " \n",
307
+ " # Skip excluded territories\n",
308
+ " if country.lower() in excluded_territories:\n",
309
+ " continue\n",
310
+ "\n",
311
+ " country_lc = country.lower()\n",
312
+ " nationality_lc = nationality.lower()\n",
313
+ "\n",
314
+ " # Store as (country, nationality, priority)\n",
315
+ " # Exact country name match = highest priority\n",
316
+ " if country_lc not in tag_to_country_nationality:\n",
317
+ " tag_to_country_nationality[country_lc] = (country, \"\", 3)\n",
318
+ " \n",
319
+ " # Exact nationality match = medium priority \n",
320
+ " if nationality_lc not in tag_to_country_nationality:\n",
321
+ " tag_to_country_nationality[nationality_lc] = (\"\", nationality, 2)\n",
322
+ " \n",
323
+ " # No-space versions\n",
324
+ " country_no_space = country_lc.replace(\" \", \"\")\n",
325
+ " nationality_no_space = nationality_lc.replace(\" \", \"\")\n",
326
+ " \n",
327
+ " if country_no_space not in tag_to_country_nationality:\n",
328
+ " tag_to_country_nationality[country_no_space] = (country, \"\", 3)\n",
329
+ " if nationality_no_space not in tag_to_country_nationality:\n",
330
+ " tag_to_country_nationality[nationality_no_space] = (\"\", nationality, 2)\n",
331
+ "\n",
332
+ " # Word parts = lowest priority (only for longer words to avoid false matches)\n",
333
+ " for part in country_lc.split():\n",
334
+ " if len(part) > 4: # Only words longer than 4 chars\n",
335
+ " if part not in tag_to_country_nationality:\n",
336
+ " tag_to_country_nationality[part] = (country, \"\", 1)\n",
337
+ " for part in nationality_lc.split():\n",
338
+ " if len(part) > 4:\n",
339
+ " if part not in tag_to_country_nationality:\n",
340
+ " tag_to_country_nationality[part] = (\"\", nationality, 1)\n",
341
+ "\n",
342
+ "print(f\"Built country/nationality mapping with {len(tag_to_country_nationality)} entries\")\n",
343
+ "\n",
344
+ "# Step 3: Infer likely_country and likely_nationality by checking ALL tags\n",
345
+ "def infer_country_and_nationality(tags):\n",
346
+ " \"\"\"\n",
347
+ " Check ALL tags and return the best match based on priority.\n",
348
+ " Priority: exact country name > nationality > word parts\n",
349
+ " \"\"\"\n",
350
+ " best_match = None\n",
351
+ " best_priority = 0\n",
352
+ " \n",
353
+ " for tag in tags:\n",
354
+ " # Try cleaned version (no spaces)\n",
355
+ " cleaned = tag.replace(\" \", \"\").lower()\n",
356
+ " \n",
357
+ " # Check cleaned version\n",
358
+ " if cleaned in tag_to_country_nationality:\n",
359
+ " country, nationality, priority = tag_to_country_nationality[cleaned]\n",
360
+ " if priority > best_priority and country and country.lower() not in excluded_territories:\n",
361
+ " best_match = (country, nationality)\n",
362
+ " best_priority = priority\n",
363
+ " \n",
364
+ " # Also check original tag\n",
365
+ " if tag in tag_to_country_nationality:\n",
366
+ " country, nationality, priority = tag_to_country_nationality[tag]\n",
367
+ " if priority > best_priority and country and country.lower() not in excluded_territories:\n",
368
+ " best_match = (country, nationality)\n",
369
+ " best_priority = priority\n",
370
+ " \n",
371
+ " if best_match:\n",
372
+ " return pd.Series(best_match)\n",
373
+ " return pd.Series([\"\", \"\"])\n",
374
+ "\n",
375
+ "poi_df[[\"likely_country\", \"likely_nationality\"]] = poi_df[\"tags\"].apply(infer_country_and_nationality)\n",
376
+ "\n",
377
+ "# Step 4: Build tag → profession mapping\n",
378
+ "profession_alias_map = {}\n",
379
+ "\n",
380
+ "for _, row in professions_df.iterrows():\n",
381
+ " canonical = str(row['profession']).strip().lower()\n",
382
+ " profession_alias_map[canonical] = canonical\n",
383
+ " for alias_col in ['alias_1', 'alias_2', 'alias_3']:\n",
384
+ " alias = row.get(alias_col)\n",
385
+ " if pd.notna(alias):\n",
386
+ " profession_alias_map[str(alias).strip().lower()] = canonical\n",
387
+ "\n",
388
+ "# Step 5: Infer likely profession from tags\n",
389
+ "def infer_profession_from_tags(tags):\n",
390
+ " matched = []\n",
391
+ " for tag in tags:\n",
392
+ " cleaned = tag.strip().lower()\n",
393
+ " if cleaned in profession_alias_map:\n",
394
+ " matched.append(profession_alias_map[cleaned])\n",
395
+ "\n",
396
+ " if not matched:\n",
397
+ " return \"\"\n",
398
+ " if \"celebrity\" in matched and len(set(matched)) > 1:\n",
399
+ " # Drop 'celebrity' if other professions are present\n",
400
+ " matched = [m for m in matched if m != \"celebrity\"]\n",
401
+ "\n",
402
+ " return matched[0] # Return the first specific match\n",
403
+ "\n",
404
+ "\n",
405
+ "poi_df[\"likely_profession\"] = poi_df[\"tags\"].apply(infer_profession_from_tags)\n",
406
+ "\n",
407
+ "# Step 6: Save enriched dataset\n",
408
+ "poi_df.to_csv(output_file, index=False)\n",
409
+ "\n",
410
+ "# Preview results\n",
411
+ "print(f\"\\nProcessed {len(poi_df)} rows\")\n",
412
+ "print(f\"Rows with country: {(poi_df['likely_country'] != '').sum()}\")\n",
413
+ "print(f\"Rows with nationality: {(poi_df['likely_nationality'] != '').sum()}\")\n",
414
+ "print(f\"Rows with profession: {(poi_df['likely_profession'] != '').sum()}\")\n",
415
+ "\n",
416
+ "print(f\"\\nTop 10 countries:\")\n",
417
+ "print(poi_df[poi_df['likely_country'] != '']['likely_country'].value_counts().head(10))\n"
418
+ ]
419
+ },
420
+ {
421
+ "cell_type": "markdown",
422
+ "id": "4a4a58b3",
423
+ "metadata": {},
424
+ "source": [
425
+ "## LLM ANNOTATION"
426
+ ]
427
+ },
428
+ {
429
+ "cell_type": "markdown",
430
+ "id": "b298844d",
431
+ "metadata": {},
432
+ "source": [
433
+ "#### Model Configurations"
434
+ ]
435
+ },
436
+ {
437
+ "cell_type": "code",
438
+ "execution_count": null,
439
+ "id": "39f3d65e",
440
+ "metadata": {},
441
+ "outputs": [],
442
+ "source": [
443
+ "import pandas as pd\n",
444
+ "import json\n",
445
+ "import time\n",
446
+ "import re\n",
447
+ "from pathlib import Path\n",
448
+ "from tqdm import tqdm\n",
449
+ "import torch\n",
450
+ "from transformers import AutoModelForCausalLM, AutoTokenizer, BitsAndBytesConfig\n",
451
+ "import signal\n",
452
+ "from contextlib import contextmanager\n",
453
+ "\n",
454
+ "# Configuration\n",
455
+ "current_dir = Path.cwd()\n",
456
+ "CACHE_DIR = current_dir.parent / \"data/models\"\n",
457
+ "CACHE_DIR.mkdir(parents=True, exist_ok=True)\n",
458
+ "\n",
459
+ "# Model configurations\n",
460
+ "MODEL_CONFIGS = {\n",
461
+ " 'mistral': {\n",
462
+ " 'name': 'mistralai/Mistral-7B-Instruct-v0.3',\n",
463
+ " 'dtype': torch.bfloat16,\n",
464
+ " 'quantization': None,\n",
465
+ " 'generation_params': {\n",
466
+ " 'max_new_tokens': 512,\n",
467
+ " 'temperature': 0.05,\n",
468
+ " 'do_sample': True,\n",
469
+ " 'top_p': 0.8,\n",
470
+ " }\n",
471
+ " },\n",
472
+ " 'gemma': {\n",
473
+ " 'name': 'google/gemma-3-27b-it',\n",
474
+ " 'dtype': torch.bfloat16,\n",
475
+ " 'quantization': None,\n",
476
+ " 'generation_params': {\n",
477
+ " 'max_new_tokens': 512,\n",
478
+ " 'temperature': 0.1,\n",
479
+ " 'do_sample': True,\n",
480
+ " 'top_p': 0.9,\n",
481
+ " }\n",
482
+ " },\n",
483
+ " 'qwen': {\n",
484
+ " 'name': 'Qwen/Qwen2.5-32B-Instruct',\n",
485
+ " 'dtype': None, # Will use quantization\n",
486
+ " 'quantization': BitsAndBytesConfig(\n",
487
+ " load_in_8bit=True,\n",
488
+ " llm_int8_threshold=6.0,\n",
489
+ " llm_int8_has_fp16_weight=False\n",
490
+ " ),\n",
491
+ " 'generation_params': {\n",
492
+ " 'max_new_tokens': 100,\n",
493
+ " 'temperature': 0.1,\n",
494
+ " 'do_sample': False,\n",
495
+ " }\n",
496
+ " }\n",
497
+ "}\n",
498
+ "\n",
499
+ "PROFESSION_CATEGORIES = [\n",
500
+ " \"actor\",\n",
501
+ " \"adult performer\",\n",
502
+ " \"singer/musician\",\n",
503
+ " \"model\",\n",
504
+ " \"online personality\",\n",
505
+ " \"public figure\",\n",
506
+ " \"voice actor/ASMR\",\n",
507
+ " \"sports professional\",\n",
508
+ " \"tv personality\"\n",
509
+ "]\n"
510
+ ]
511
+ },
512
+ {
513
+ "cell_type": "markdown",
514
+ "id": "c215b38c",
515
+ "metadata": {},
516
+ "source": [
517
+ "#### Load Model Function"
518
+ ]
519
+ },
520
+ {
521
+ "cell_type": "code",
522
+ "execution_count": null,
523
+ "id": "cfb5b13e",
524
+ "metadata": {},
525
+ "outputs": [],
526
+ "source": [
527
+ "def load_model(model_type='mistral'):\n",
528
+ " \"\"\"\n",
529
+ " Load model and tokenizer based on type.\n",
530
+ " \n",
531
+ " Args:\n",
532
+ " model_type: 'mistral', 'gemma', or 'qwen'\n",
533
+ " \n",
534
+ " Returns:\n",
535
+ " tuple: (model, tokenizer, config)\n",
536
+ " \"\"\"\n",
537
+ " if model_type not in MODEL_CONFIGS:\n",
538
+ " raise ValueError(f\"Unknown model type: {model_type}. Choose from {list(MODEL_CONFIGS.keys())}\")\n",
539
+ " \n",
540
+ " config = MODEL_CONFIGS[model_type]\n",
541
+ " model_name = config['name']\n",
542
+ " \n",
543
+ " device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n",
544
+ " print(f\"Loading model: {model_name}\")\n",
545
+ " print(f\"Cache directory: {CACHE_DIR}\")\n",
546
+ " print(f\"Device: {device}\\n\")\n",
547
+ " \n",
548
+ " if device == \"cpu\":\n",
549
+ " print(\"⚠️ WARNING: No GPU detected! Inference will be VERY slow.\")\n",
550
+ " \n",
551
+ " # Load tokenizer\n",
552
+ " try:\n",
553
+ " tokenizer = AutoTokenizer.from_pretrained(\n",
554
+ " model_name,\n",
555
+ " cache_dir=str(CACHE_DIR),\n",
556
+ " use_fast=True\n",
557
+ " )\n",
558
+ " except:\n",
559
+ " tokenizer = AutoTokenizer.from_pretrained(\n",
560
+ " model_name,\n",
561
+ " cache_dir=str(CACHE_DIR),\n",
562
+ " use_fast=False\n",
563
+ " )\n",
564
+ " \n",
565
+ " if tokenizer.pad_token is None:\n",
566
+ " tokenizer.pad_token = tokenizer.eos_token\n",
567
+ " \n",
568
+ " # Load model\n",
569
+ " model_kwargs = {\n",
570
+ " 'cache_dir': str(CACHE_DIR),\n",
571
+ " 'device_map': 'auto',\n",
572
+ " 'trust_remote_code': False\n",
573
+ " }\n",
574
+ " \n",
575
+ " if config['quantization']:\n",
576
+ " model_kwargs['quantization_config'] = config['quantization']\n",
577
+ " else:\n",
578
+ " model_kwargs['torch_dtype'] = config['dtype']\n",
579
+ " \n",
580
+ " model = AutoModelForCausalLM.from_pretrained(model_name, **model_kwargs)\n",
581
+ " model.eval()\n",
582
+ " \n",
583
+ " # Check VRAM\n",
584
+ " if torch.cuda.is_available():\n",
585
+ " vram_gb = torch.cuda.max_memory_allocated() / 1024**3\n",
586
+ " print(f\"VRAM used: {vram_gb:.2f} GB\\n\")\n",
587
+ " \n",
588
+ " return model, tokenizer, config\n"
589
+ ]
590
+ },
591
+ {
592
+ "cell_type": "markdown",
593
+ "id": "11b2221a",
594
+ "metadata": {},
595
+ "source": [
596
+ "#### Inference Code"
597
+ ]
598
+ },
599
+ {
600
+ "cell_type": "code",
601
+ "execution_count": null,
602
+ "id": "229f96bd",
603
+ "metadata": {},
604
+ "outputs": [],
605
+ "source": [
606
+ "@contextmanager\n",
607
+ "def timeout(duration):\n",
608
+ " \"\"\"Context manager for timeout.\"\"\"\n",
609
+ " def handler(signum, frame):\n",
610
+ " raise TimeoutError(\"Operation timed out\")\n",
611
+ " \n",
612
+ " signal.signal(signal.SIGALRM, handler)\n",
613
+ " signal.alarm(duration)\n",
614
+ " try:\n",
615
+ " yield\n",
616
+ " finally:\n",
617
+ " signal.alarm(0)\n",
618
+ "\n",
619
+ "def query_model(prompt, model, tokenizer, config, use_timeout=False):\n",
620
+ " \"\"\"\n",
621
+ " Query model with given prompt.\n",
622
+ " \n",
623
+ " Args:\n",
624
+ " prompt: Input prompt string\n",
625
+ " model: Loaded model\n",
626
+ " tokenizer: Loaded tokenizer\n",
627
+ " config: Model configuration dict\n",
628
+ " use_timeout: Whether to use 60s timeout (for Qwen)\n",
629
+ " \n",
630
+ " Returns:\n",
631
+ " str: Model response or None on error\n",
632
+ " \"\"\"\n",
633
+ " try:\n",
634
+ " device = next(model.parameters()).device\n",
635
+ " \n",
636
+ " # Format as chat message\n",
637
+ " messages = [\n",
638
+ " {\"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",
639
+ " {\"role\": \"user\", \"content\": prompt}\n",
640
+ " ]\n",
641
+ " \n",
642
+ " # Tokenize\n",
643
+ " if hasattr(tokenizer, 'apply_chat_template'):\n",
644
+ " text = tokenizer.apply_chat_template(\n",
645
+ " messages,\n",
646
+ " tokenize=False,\n",
647
+ " add_generation_prompt=True\n",
648
+ " )\n",
649
+ " else:\n",
650
+ " text = f\"[INST] {prompt} [/INST]\"\n",
651
+ " \n",
652
+ " inputs = tokenizer([text], return_tensors=\"pt\", padding=True).to(device)\n",
653
+ " \n",
654
+ " # Generation parameters\n",
655
+ " gen_kwargs = config['generation_params'].copy()\n",
656
+ " gen_kwargs['pad_token_id'] = tokenizer.eos_token_id\n",
657
+ " \n",
658
+ " # Generate\n",
659
+ " generation_fn = lambda: model.generate(**inputs, **gen_kwargs)\n",
660
+ " \n",
661
+ " if use_timeout:\n",
662
+ " with timeout(60):\n",
663
+ " with torch.no_grad():\n",
664
+ " outputs = generation_fn()\n",
665
+ " else:\n",
666
+ " with torch.no_grad():\n",
667
+ " outputs = generation_fn()\n",
668
+ " \n",
669
+ " # Decode\n",
670
+ " generated_ids = outputs[0][inputs['input_ids'].shape[1]:]\n",
671
+ " response = tokenizer.decode(generated_ids, skip_special_tokens=True)\n",
672
+ " \n",
673
+ " return response.strip()\n",
674
+ " \n",
675
+ " except TimeoutError:\n",
676
+ " print(f\"[ERROR] Generation timed out after 60 seconds\")\n",
677
+ " return None\n",
678
+ " except Exception as e:\n",
679
+ " print(f\"[ERROR] Generation failed: {e}\")\n",
680
+ " return None\n"
681
+ ]
682
+ },
683
+ {
684
+ "cell_type": "markdown",
685
+ "id": "88f005f8",
686
+ "metadata": {},
687
+ "source": [
688
+ "#### Prompt creation"
689
+ ]
690
+ },
691
+ {
692
+ "cell_type": "code",
693
+ "execution_count": null,
694
+ "id": "dfe05463",
695
+ "metadata": {},
696
+ "outputs": [],
697
+ "source": [
698
+ "def create_prompt(row):\n",
699
+ " \"\"\"Create annotation prompt from row data.\"\"\"\n",
700
+ " name = row['real_name'] if pd.notna(row.get('real_name')) else row.get('name', '')\n",
701
+ " \n",
702
+ " # Gather hints\n",
703
+ " hints = []\n",
704
+ " if pd.notna(row.get('likely_profession')):\n",
705
+ " hints.append(str(row['likely_profession']))\n",
706
+ " if pd.notna(row.get('likely_nationality')):\n",
707
+ " hints.append(str(row['likely_nationality']))\n",
708
+ " if pd.notna(row.get('likely_country')):\n",
709
+ " hints.append(str(row['likely_country']))\n",
710
+ " \n",
711
+ " # Add tags if needed\n",
712
+ " if len(hints) < 3:\n",
713
+ " for i in range(1, 8):\n",
714
+ " tag_col = f'tag_{i}'\n",
715
+ " if tag_col in row and pd.notna(row[tag_col]):\n",
716
+ " tag_val = str(row[tag_col])\n",
717
+ " if tag_val not in hints:\n",
718
+ " hints.append(tag_val)\n",
719
+ " if len(hints) >= 5:\n",
720
+ " break\n",
721
+ " \n",
722
+ " hint_text = \", \".join(hints[:5]) if hints else \"none\"\n",
723
+ " \n",
724
+ " return f\"\"\"Extract information about '{name}' ({hint_text}).\n",
725
+ "\n",
726
+ "Respond with EXACTLY 5 numbered lines. Each line must contain ONLY the value, no labels or extra text.\n",
727
+ "\n",
728
+ "FORMAT REQUIREMENTS:\n",
729
+ "1. Full legal name in Western order (first last). VALUE ONLY.\n",
730
+ "2. Stage names/aliases, comma-separated. If none, write \"None\". VALUE ONLY.\n",
731
+ "3. Gender: MUST be exactly one word: Male, Female, Other, or Unknown. VALUE ONLY.\n",
732
+ "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",
733
+ "5. Primary country: Country name only (e.g., \"China\", \"United States\", \"Colombia\"). VALUE ONLY.\n",
734
+ "\n",
735
+ "RULES:\n",
736
+ "- Professions MUST match the exact categories listed (actress = actor)\n",
737
+ "- \"online personality\" includes streamers, cosplayers, YouTubers, influencers\n",
738
+ "- \"public figure\" includes politicians, activists, journalists, authors\n",
739
+ "- Use \"Unknown\" when uncertain or for fictional characters\n",
740
+ "- NO explanatory text, NO labels like \"Gender:\", NO prefixes\n",
741
+ "- For multi-role people, list up to 3 categories by relevance\n",
742
+ "\n",
743
+ "EXAMPLE FORMAT:\n",
744
+ "1. Taylor Swift\n",
745
+ "2. None\n",
746
+ "3. Female\n",
747
+ "4. singer/musician, public figure\n",
748
+ "5. United States\"\"\"\n"
749
+ ]
750
+ },
751
+ {
752
+ "cell_type": "markdown",
753
+ "id": "854fa668",
754
+ "metadata": {},
755
+ "source": [
756
+ "#### Response parsing code"
757
+ ]
758
+ },
759
+ {
760
+ "cell_type": "code",
761
+ "execution_count": null,
762
+ "id": "1a4be2ee",
763
+ "metadata": {},
764
+ "outputs": [],
765
+ "source": [
766
+ "def parse_response(response):\n",
767
+ " \"\"\"Parse model response into structured fields.\"\"\"\n",
768
+ " if not response:\n",
769
+ " return {\n",
770
+ " 'full_name': 'Unknown',\n",
771
+ " 'aliases': 'Unknown',\n",
772
+ " 'gender': 'Unknown',\n",
773
+ " 'profession_llm': 'Unknown',\n",
774
+ " 'country': 'Unknown'\n",
775
+ " }\n",
776
+ " \n",
777
+ " lines = [line.strip() for line in response.split('\\n') if line.strip()]\n",
778
+ " \n",
779
+ " fields = {\n",
780
+ " 'full_name': 'Unknown',\n",
781
+ " 'aliases': 'Unknown',\n",
782
+ " 'gender': 'Unknown',\n",
783
+ " 'profession_llm': 'Unknown',\n",
784
+ " 'country': 'Unknown'\n",
785
+ " }\n",
786
+ " \n",
787
+ " for line in lines:\n",
788
+ " if line.startswith('1.'):\n",
789
+ " fields['full_name'] = line[2:].strip()\n",
790
+ " elif line.startswith('2.'):\n",
791
+ " fields['aliases'] = line[2:].strip()\n",
792
+ " elif line.startswith('3.'):\n",
793
+ " gender_raw = line[2:].strip()\n",
794
+ " gender_raw = re.sub(r'^(Gender:|gender:)\\s*', '', gender_raw, flags=re.IGNORECASE)\n",
795
+ " gender_match = re.search(r'\\b(Male|Female|Other|Unknown)\\b', gender_raw, re.IGNORECASE)\n",
796
+ " fields['gender'] = gender_match.group(1).capitalize() if gender_match else gender_raw\n",
797
+ " elif line.startswith('4.'):\n",
798
+ " fields['profession_llm'] = line[2:].strip()\n",
799
+ " elif line.startswith('5.'):\n",
800
+ " country_raw = line[2:].strip()\n",
801
+ " country_raw = re.sub(r'^(Primary\\s+)?(associated\\s+)?country:\\s*', '', country_raw, flags=re.IGNORECASE)\n",
802
+ " fields['country'] = country_raw\n",
803
+ " \n",
804
+ " return fields\n"
805
+ ]
806
+ },
807
+ {
808
+ "cell_type": "markdown",
809
+ "id": "7e2f7a86",
810
+ "metadata": {},
811
+ "source": [
812
+ "#### CSV annotation"
813
+ ]
814
+ },
815
+ {
816
+ "cell_type": "code",
817
+ "execution_count": null,
818
+ "id": "5f3dd5d6",
819
+ "metadata": {},
820
+ "outputs": [],
821
+ "source": [
822
+ "def annotate_dataset(model_type='mistral', test_mode=False, test_size=100, max_rows=50862, save_interval=10):\n",
823
+ " \"\"\"\n",
824
+ " Annotate dataset using specified model.\n",
825
+ " \n",
826
+ " Args:\n",
827
+ " model_type: 'mistral', 'gemma', or 'qwen'\n",
828
+ " test_mode: If True, only process test_size rows\n",
829
+ " test_size: Number of rows to process in test mode\n",
830
+ " max_rows: Maximum rows to process\n",
831
+ " save_interval: Save progress every N rows\n",
832
+ " \"\"\"\n",
833
+ " # Setup paths\n",
834
+ " input_file = current_dir.parent / \"data/CSV/model_adapter/real_person_adapter_step_02_NER.csv\"\n",
835
+ " output_file = current_dir.parent / f\"data/CSV/{model_type}_local_annotated_POI{'_test' if test_mode else ''}.csv\"\n",
836
+ " index_file = current_dir.parent / f\"misc/query_indicies/{model_type}_local_query_index.txt\"\n",
837
+ " index_file.parent.mkdir(parents=True, exist_ok=True)\n",
838
+ " \n",
839
+ " # Load model\n",
840
+ " model, tokenizer, config = load_model(model_type)\n",
841
+ " \n",
842
+ " # Load data\n",
843
+ " print(f\"Loaded {len(df)} rows from input file\")\n",
844
+ " df = pd.read_csv(input_file)\n",
845
+ " \n",
846
+ " # Merge existing annotations if available\n",
847
+ " if output_file.exists():\n",
848
+ " existing_df = pd.read_csv(output_file)\n",
849
+ " annotation_cols = ['full_name', 'aliases', 'gender', 'profession_llm', 'country']\n",
850
+ " for col in annotation_cols:\n",
851
+ " if col in existing_df.columns:\n",
852
+ " df[col] = existing_df[col][:len(df)]\n",
853
+ " \n",
854
+ " # Apply limits\n",
855
+ " if test_mode:\n",
856
+ " df = df.head(test_size).copy()\n",
857
+ " elif max_rows:\n",
858
+ " df = df.head(max_rows).copy()\n",
859
+ " \n",
860
+ " # Create prompts\n",
861
+ " df['prompt'] = df.apply(create_prompt, axis=1)\n",
862
+ " \n",
863
+ " # Load progress index\n",
864
+ " current_index = 0\n",
865
+ " if index_file.exists():\n",
866
+ " try:\n",
867
+ " current_index = int(index_file.read_text().strip())\n",
868
+ " except:\n",
869
+ " current_index = 0\n",
870
+ " \n",
871
+ " print(f\"Resuming from index {current_index}\")\n",
872
+ " \n",
873
+ " # Process rows\n",
874
+ " use_timeout = (model_type == 'qwen')\n",
875
+ " \n",
876
+ " for i in tqdm(range(current_index, len(df)), desc=f\"{model_type.capitalize()} Annotation\"):\n",
877
+ " prompt = df.at[i, \"prompt\"]\n",
878
+ " \n",
879
+ " # Query with retries\n",
880
+ " response = None\n",
881
+ " for attempt in range(3):\n",
882
+ " response = query_model(prompt, model, tokenizer, config, use_timeout)\n",
883
+ " \n",
884
+ " if response and len(response.strip()) > 10:\n",
885
+ " break\n",
886
+ " \n",
887
+ " print(f\"⚠️ Row {i}: Empty or invalid response, retry {attempt+1}/3\")\n",
888
+ " time.sleep(0.5)\n",
889
+ " \n",
890
+ " # Skip if invalid\n",
891
+ " if not response or len(response.strip()) <= 10:\n",
892
+ " print(f\"❌ Row {i}: failed after retries, skipping\")\n",
893
+ " continue\n",
894
+ " \n",
895
+ " # Parse and validate\n",
896
+ " parsed = parse_response(response)\n",
897
+ " \n",
898
+ " if all(v == \"Unknown\" for v in parsed.values()):\n",
899
+ " print(f\"❌ Row {i}: parsed as all Unknown, skipping\")\n",
900
+ " continue\n",
901
+ " \n",
902
+ " # Write fields\n",
903
+ " for key, value in parsed.items():\n",
904
+ " df.at[i, key] = value\n",
905
+ " \n",
906
+ " current_index = i + 1\n",
907
+ " \n",
908
+ " # GPU cleanup\n",
909
+ " if torch.cuda.is_available():\n",
910
+ " torch.cuda.empty_cache()\n",
911
+ " torch.cuda.synchronize()\n",
912
+ " \n",
913
+ " # Save progress\n",
914
+ " if (i + 1) % save_interval == 0 or (i + 1) == len(df):\n",
915
+ " df.to_csv(output_file, index=False)\n",
916
+ " index_file.write_text(str(current_index))\n",
917
+ " print(f\"💾 Progress saved after row {i+1}\")\n",
918
+ " \n",
919
+ " # Final save\n",
920
+ " df.to_csv(output_file, index=False)\n",
921
+ " index_file.write_text(str(current_index))\n",
922
+ " print(f\"✓ Finished annotation with {model_type}\")\n"
923
+ ]
924
+ },
925
+ {
926
+ "cell_type": "markdown",
927
+ "id": "55da2f4c",
928
+ "metadata": {},
929
+ "source": [
930
+ "### Usage Examples\n",
931
+ "Run annotation with your chosen model."
932
+ ]
933
+ },
934
+ {
935
+ "cell_type": "code",
936
+ "execution_count": null,
937
+ "id": "351ea40c",
938
+ "metadata": {},
939
+ "outputs": [],
940
+ "source": [
941
+ "# Example 1: Annotate with Mistral (13.5 GB VRAM)\n",
942
+ "# annotate_dataset(model_type='mistral', test_mode=False)\n",
943
+ "\n",
944
+ "# Example 2: Annotate with Gemma (56.3 GB VRAM)\n",
945
+ "# annotate_dataset(model_type='gemma', test_mode=False)\n",
946
+ "\n",
947
+ "# Example 3: Annotate with Qwen (32.7 GB VRAM, 8-bit)\n",
948
+ "# annotate_dataset(model_type='qwen', test_mode=False)\n",
949
+ "\n",
950
+ "# Test mode (first 100 rows)\n",
951
+ "# annotate_dataset(model_type='mistral', test_mode=True, test_size=100)\n"
952
+ ]
953
+ },
954
+ {
955
+ "cell_type": "code",
956
+ "execution_count": null,
957
+ "id": "e8203abc-e7c3-4cb6-aaeb-fdc6933981fc",
958
+ "metadata": {},
959
+ "outputs": [],
960
+ "source": [
961
+ "import pandas as pd\n",
962
+ "import json\n",
963
+ "import time\n",
964
+ "import re\n",
965
+ "from pathlib import Path\n",
966
+ "from tqdm import tqdm\n",
967
+ "import torch\n",
968
+ "from transformers import AutoModelForCausalLM, AutoTokenizer, BitsAndBytesConfig\n",
969
+ "import signal\n",
970
+ "from contextlib import contextmanager\n",
971
+ "\n",
972
+ "current_dir = Path.cwd()\n",
973
+ "input_file = current_dir.parent / \"data/CSV/model_adapter/real_person_adapter_step_02_NER.csv\"\n",
974
+ "professions_file = current_dir.parent / \"misc/lists/professions.csv\"\n",
975
+ "professions_mapped_file = current_dir.parent / \"misc/lists/professions_mapped.csv\"\n",
976
+ "# === PROCESS DATA ===\n",
977
+ "\n",
978
+ "\n",
979
+ "# === CONFIGURATION ===\n",
980
+ "TEST_MODE = False\n",
981
+ "TEST_SIZE = 100\n",
982
+ "MAX_ROWS = 50862\n",
983
+ "SAVE_INTERVAL = 10\n",
984
+ "\n",
985
+ "\n",
986
+ "index_file = current_dir.parent / \"misc/query_indicies/eurollm_local_query_index.txt\"\n",
987
+ "output_file = current_dir.parent / f\"data/CSV/eurollm_local_annotated_POI{'_test' if TEST_MODE else ''}.csv\"\n",
988
+ "\n",
989
+ "# Model settings\n",
990
+ "MODEL_NAME = \"utter-project/EuroLLM-9B\"\n",
991
+ "#MODEL_NAME = \"Qwen/Qwen2.5-32B-Instruct\"\n",
992
+ "#MODEL_NAME = \"Qwen/Qwen2.5-14B-Instruct\"\n",
993
+ "#MODEL_NAME = \"Qwen/Qwen3-235B-A22B-Instruct-2507-FP8\"\n",
994
+ "#MODEL_NAME = \"mistralai/Mistral-7B-Instruct-v0.3\"\n",
995
+ "CACHE_DIR = current_dir.parent / \"data/models\"\n",
996
+ "CACHE_DIR.mkdir(parents=True, exist_ok=True)\n",
997
+ "\n",
998
+ "# Define the SPECIFIC profession categories\n",
999
+ "PROFESSION_CATEGORIES = [\n",
1000
+ " \"actor\",\n",
1001
+ " \"adult performer\",\n",
1002
+ " \"singer/musician\",\n",
1003
+ " \"model\",\n",
1004
+ " \"online personality\",\n",
1005
+ " \"public figure\",\n",
1006
+ " \"voice actor/ASMR\",\n",
1007
+ " \"sports professional\",\n",
1008
+ " \"tv personality\"\n",
1009
+ "]\n",
1010
+ "\n",
1011
+ "# === LOAD MODEL ===\n",
1012
+ "print(f\"Loading model: {MODEL_NAME}\")\n",
1013
+ "print(f\"Cache directory: {CACHE_DIR}\")\n",
1014
+ "print(f\"This may take a while on first run...\\n\")\n",
1015
+ "\n",
1016
+ "# Check GPU availability\n",
1017
+ "device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n",
1018
+ "print(f\"Device: {device}\")\n",
1019
+ "\n",
1020
+ "if device == \"cpu\":\n",
1021
+ " print(\"⚠️ WARNING: No GPU detected! Inference will be VERY slow.\")\n",
1022
+ " print(\" Consider using a GPU or reducing model size.\")\n",
1023
+ "\n",
1024
+ "# Get HF token from credentials file\n",
1025
+ "import os\n",
1026
+ "credentials_dir = current_dir.parent / \"misc/credentials\"\n",
1027
+ "hf_token_file = credentials_dir / \"hf_token.txt\"\n",
1028
+ "\n",
1029
+ "HF_TOKEN = None\n",
1030
+ "if hf_token_file.exists():\n",
1031
+ " HF_TOKEN = hf_token_file.read_text().strip()\n",
1032
+ " print(\"✅ HF token loaded from credentials file\")\n",
1033
+ "else:\n",
1034
+ " print(\"⚠️ HF token file not found at:\", hf_token_file)\n",
1035
+ " print(\" The script will try to use cached credentials from 'huggingface-cli login'\")\n",
1036
+ " print(\" Or create the file: misc/credentials/hf_token.txt with your token\")\n",
1037
+ " HF_TOKEN = None # Will use cached token if available\n",
1038
+ "\n",
1039
+ "# Load tokenizer\n",
1040
+ "print(\"Loading tokenizer...\")\n",
1041
+ "try:\n",
1042
+ " tokenizer = AutoTokenizer.from_pretrained(\n",
1043
+ " MODEL_NAME,\n",
1044
+ " cache_dir=str(CACHE_DIR),\n",
1045
+ " use_fast=True,\n",
1046
+ " token=HF_TOKEN\n",
1047
+ " )\n",
1048
+ "except Exception as e:\n",
1049
+ " print(f\"Failed with use_fast=True, trying use_fast=False...\")\n",
1050
+ " tokenizer = AutoTokenizer.from_pretrained(\n",
1051
+ " MODEL_NAME,\n",
1052
+ " cache_dir=str(CACHE_DIR),\n",
1053
+ " use_fast=False,\n",
1054
+ " token=HF_TOKEN\n",
1055
+ " )\n",
1056
+ "\n",
1057
+ "# Ensure pad token is set\n",
1058
+ "if tokenizer.pad_token is None:\n",
1059
+ " tokenizer.pad_token = tokenizer.eos_token\n",
1060
+ "\n",
1061
+ "print(\"✅ Tokenizer loaded\")\n",
1062
+ "\n",
1063
+ "# Configure 8-bit quantization for A100\n",
1064
+ "print(\"Configuring 8-bit quantization...\")\n",
1065
+ "quantization_config = BitsAndBytesConfig(\n",
1066
+ " load_in_8bit=True,\n",
1067
+ " llm_int8_threshold=6.0,\n",
1068
+ " llm_int8_has_fp16_weight=False\n",
1069
+ ")\n",
1070
+ "\n",
1071
+ "# Load model with 8-bit quantization\n",
1072
+ "print(\"Loading model with 8-bit quantization (this may take several minutes)...\")\n",
1073
+ "model = AutoModelForCausalLM.from_pretrained(\n",
1074
+ " MODEL_NAME,\n",
1075
+ " cache_dir=str(CACHE_DIR),\n",
1076
+ " quantization_config=quantization_config,\n",
1077
+ " device_map=\"auto\",\n",
1078
+ " trust_remote_code=False,\n",
1079
+ " token=HF_TOKEN\n",
1080
+ ")\n",
1081
+ "model.eval()\n",
1082
+ "print(\"✅ Model loaded with 8-bit quantization\")\n",
1083
+ "\n",
1084
+ "# Check VRAM usage\n",
1085
+ "if torch.cuda.is_available():\n",
1086
+ " vram_gb = torch.cuda.max_memory_allocated() / 1024**3\n",
1087
+ " print(f\"VRAM used: {vram_gb:.2f} GB\\n\")\n",
1088
+ "\n",
1089
+ "# === LOAD DATA ===\n",
1090
+ "if output_file.exists():\n",
1091
+ " print(\"Loading annotated CSV...\")\n",
1092
+ " df = pd.read_csv(output_file)\n",
1093
+ "else:\n",
1094
+ " print(\"Loading raw input CSV...\")\n",
1095
+ " df = pd.read_csv(input_file)\n",
1096
+ "\n",
1097
+ "\n",
1098
+ "# Try to load profession mapping files\n",
1099
+ "try:\n",
1100
+ " professions_df = pd.read_csv(professions_file)\n",
1101
+ " print(f\"✅ Loaded professions.csv\")\n",
1102
+ "except:\n",
1103
+ " print(\"⚠️ Warning: professions.csv not found\")\n",
1104
+ "\n",
1105
+ "try:\n",
1106
+ " prof_mapped_df = pd.read_csv(professions_mapped_file)\n",
1107
+ " print(f\"✅ Loaded profession mapping with {len(prof_mapped_df)} categories\")\n",
1108
+ "except:\n",
1109
+ " print(\"⚠️ Warning: professions_mapped.csv not found, using default categories\")\n",
1110
+ "\n",
1111
+ "profession_str = \", \".join(PROFESSION_CATEGORIES)\n",
1112
+ "\n",
1113
+ "print(f\"Loaded {len(df)} rows\")\n",
1114
+ "print(f\"\\nProfession categories ({len(PROFESSION_CATEGORIES)}):\")\n",
1115
+ "for cat in PROFESSION_CATEGORIES:\n",
1116
+ " print(f\" - {cat}\")\n",
1117
+ "\n",
1118
+ "if TEST_MODE:\n",
1119
+ " print(f\"\\nRunning in TEST MODE with {TEST_SIZE} samples\")\n",
1120
+ " df = df.head(TEST_SIZE).copy()\n",
1121
+ "elif MAX_ROWS:\n",
1122
+ " df = df.head(MAX_ROWS).copy()\n",
1123
+ "\n",
1124
+ "# === CREATE PROMPTS (OPTIMIZED FOR CLEAN OUTPUTS) ===\n",
1125
+ "def create_prompt(row):\n",
1126
+ " \"\"\"Create prompt for EuroLLM annotation with strict formatting requirements.\"\"\"\n",
1127
+ " name = row['real_name'] if pd.notna(row.get('real_name')) else row.get('name', '')\n",
1128
+ " \n",
1129
+ " # Gather hints\n",
1130
+ " hints = []\n",
1131
+ " if pd.notna(row.get('likely_profession')):\n",
1132
+ " hints.append(str(row['likely_profession']))\n",
1133
+ " if pd.notna(row.get('likely_nationality')):\n",
1134
+ " hints.append(str(row['likely_nationality']))\n",
1135
+ " if pd.notna(row.get('likely_country')):\n",
1136
+ " hints.append(str(row['likely_country']))\n",
1137
+ " \n",
1138
+ " # Add tags if we don't have enough hints\n",
1139
+ " if len(hints) < 3:\n",
1140
+ " for i in range(1, 8):\n",
1141
+ " tag_col = f'tag_{i}'\n",
1142
+ " if tag_col in row and pd.notna(row[tag_col]):\n",
1143
+ " tag_val = str(row[tag_col])\n",
1144
+ " if tag_val not in hints:\n",
1145
+ " hints.append(tag_val)\n",
1146
+ " if len(hints) >= 5:\n",
1147
+ " break\n",
1148
+ " \n",
1149
+ " hint_text = \", \".join(hints[:5]) if hints else \"none\"\n",
1150
+ " \n",
1151
+ " return f\"\"\"Extract information about '{name}'. \n",
1152
+ "Context hints (DO NOT copy these as professions): {hint_text}\n",
1153
+ "\n",
1154
+ "Respond with EXACTLY 5 numbered lines. Each line must contain ONLY the value, no labels or extra text.\n",
1155
+ "\n",
1156
+ "FORMAT REQUIREMENTS:\n",
1157
+ "1. Full legal name in Western order (first last). VALUE ONLY.\n",
1158
+ "2. Stage names/aliases, comma-separated. If none, write \"None\". VALUE ONLY.\n",
1159
+ "3. Gender: MUST be exactly one word: Male, Female, Other, or Unknown. VALUE ONLY.\n",
1160
+ "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",
1161
+ "5. Primary country: Country name only (e.g., \"China\", \"United States\", \"Colombia\"). VALUE ONLY.\n",
1162
+ "6. If uncertain about an item, write \"Unknown\"\n",
1163
+ "\n",
1164
+ "CRITICAL RULES FOR PROFESSIONS (Line 4):\n",
1165
+ "- ONLY use the exact profession categories listed above\n",
1166
+ "- DO NOT use descriptive words like \"sexy\", \"photorealistic\", \"celebrity\"\n",
1167
+ "- DO NOT copy the hint words as professions\n",
1168
+ "- If uncertain write \"Unknown\"\n",
1169
+ "- Valid professions are ONLY: actor, adult performer, singer/musician, model, online personality, public figure, voice actor/ASMR, sports professional, tv personality\n",
1170
+ "- Actress = actor, streamer = online personality, YouTuber = online personality\n",
1171
+ "\n",
1172
+ "OTHER RULES:\n",
1173
+ "- Use \"Unknown\" when uncertain or for fictional characters\n",
1174
+ "- NO explanatory text, NO labels like \"Gender:\", NO prefixes\n",
1175
+ "- For multi-role people, list up to 3 categories by relevance\n",
1176
+ "\n",
1177
+ "EXAMPLE FORMAT:\n",
1178
+ "1. Taylor Swift\n",
1179
+ "2. None\n",
1180
+ "3. Female\n",
1181
+ "4. singer/musician, public figure\n",
1182
+ "5. United States\"\"\"\n",
1183
+ "\n",
1184
+ "# Create prompts\n",
1185
+ "print(\"\\nCreating prompts...\")\n",
1186
+ "df['prompt'] = df.apply(create_prompt, axis=1)\n",
1187
+ "print(\"✅ Prompts created\")\n",
1188
+ "\n",
1189
+ "@contextmanager\n",
1190
+ "def timeout(duration):\n",
1191
+ " def handler(signum, frame):\n",
1192
+ " raise TimeoutError(\"Operation timed out\")\n",
1193
+ " \n",
1194
+ " # Set the signal handler and alarm\n",
1195
+ " signal.signal(signal.SIGALRM, handler)\n",
1196
+ " signal.alarm(duration)\n",
1197
+ " try:\n",
1198
+ " yield\n",
1199
+ " finally:\n",
1200
+ " signal.alarm(0) # Disable the alarm\n",
1201
+ "\n",
1202
+ "\n",
1203
+ "def query_eurollm_local(prompt: str) -> str:\n",
1204
+ " \"\"\"Query EuroLLM locally via transformers with very low temperature.\"\"\"\n",
1205
+ " try:\n",
1206
+ " # Format as chat message for EuroLLM with strict system prompt\n",
1207
+ " messages = [\n",
1208
+ " {\"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",
1209
+ " {\"role\": \"user\", \"content\": prompt}\n",
1210
+ " ]\n",
1211
+ " \n",
1212
+ " # Tokenize\n",
1213
+ " if hasattr(tokenizer, 'apply_chat_template') and tokenizer.chat_template is not None:\n",
1214
+ " text = tokenizer.apply_chat_template(\n",
1215
+ " messages,\n",
1216
+ " tokenize=False,\n",
1217
+ " add_generation_prompt=True\n",
1218
+ " )\n",
1219
+ " else:\n",
1220
+ " # Fallback for models without chat template\n",
1221
+ " text = f\"[INST] {prompt} [/INST]\"\n",
1222
+ " \n",
1223
+ " inputs = tokenizer([text], return_tensors=\"pt\", padding=True).to(device)\n",
1224
+ " \n",
1225
+ " # Generate with timeout and very low temperature\n",
1226
+ " with timeout(60):\n",
1227
+ " with torch.no_grad():\n",
1228
+ " outputs = model.generate(\n",
1229
+ " **inputs,\n",
1230
+ " max_new_tokens=100,\n",
1231
+ " temperature=0.01, # Very low temperature for more deterministic outputs\n",
1232
+ " do_sample=True, # Must be True when temperature is set\n",
1233
+ " pad_token_id=tokenizer.eos_token_id\n",
1234
+ " )\n",
1235
+ " \n",
1236
+ " # Decode\n",
1237
+ " generated_ids = outputs[0][inputs['input_ids'].shape[1]:]\n",
1238
+ " response = tokenizer.decode(generated_ids, skip_special_tokens=True)\n",
1239
+ " \n",
1240
+ " return response.strip()\n",
1241
+ " \n",
1242
+ " except TimeoutError:\n",
1243
+ " print(f\"[ERROR] Generation timed out after 60 seconds\")\n",
1244
+ " return None\n",
1245
+ " except Exception as e:\n",
1246
+ " print(f\"Generation error: {e}\")\n",
1247
+ " import traceback\n",
1248
+ " traceback.print_exc()\n",
1249
+ " return None\n",
1250
+ "\n",
1251
+ " \n",
1252
+ "# === PARSE RESPONSE WITH CLEANING ===\n",
1253
+ "def parse_response(response):\n",
1254
+ " \"\"\"Parse EuroLLM response into structured fields with cleaning.\"\"\"\n",
1255
+ " if not response:\n",
1256
+ " return {\n",
1257
+ " 'full_name': 'Unknown',\n",
1258
+ " 'aliases': 'Unknown',\n",
1259
+ " 'gender': 'Unknown',\n",
1260
+ " 'profession_llm': 'Unknown',\n",
1261
+ " 'country': 'Unknown'\n",
1262
+ " }\n",
1263
+ " \n",
1264
+ " # Valid profession categories\n",
1265
+ " VALID_PROFESSIONS = {\n",
1266
+ " \"actor\", \"adult performer\", \"singer/musician\", \"model\", \n",
1267
+ " \"online personality\", \"public figure\", \"voice actor/asmr\", \n",
1268
+ " \"sports professional\", \"tv personality\"\n",
1269
+ " }\n",
1270
+ " \n",
1271
+ " # Split into lines and clean\n",
1272
+ " lines = [line.strip() for line in response.split('\\n') if line.strip()]\n",
1273
+ " \n",
1274
+ " # Initialize with Unknown values\n",
1275
+ " fields = {\n",
1276
+ " 'full_name': 'Unknown',\n",
1277
+ " 'aliases': 'Unknown',\n",
1278
+ " 'gender': 'Unknown',\n",
1279
+ " 'profession_llm': 'Unknown',\n",
1280
+ " 'country': 'Unknown'\n",
1281
+ " }\n",
1282
+ " \n",
1283
+ " # Extract information from each numbered line\n",
1284
+ " for line in lines:\n",
1285
+ " if line.startswith('1.'):\n",
1286
+ " fields['full_name'] = line[2:].strip()\n",
1287
+ " elif line.startswith('2.'):\n",
1288
+ " fields['aliases'] = line[2:].strip()\n",
1289
+ " elif line.startswith('3.'):\n",
1290
+ " # Clean gender field - remove any labels\n",
1291
+ " gender_raw = line[2:].strip()\n",
1292
+ " # Remove common prefixes\n",
1293
+ " gender_raw = re.sub(r'^(Gender:|gender:)\\s*', '', gender_raw, flags=re.IGNORECASE)\n",
1294
+ " # Extract just the gender word\n",
1295
+ " gender_match = re.search(r'\\b(Male|Female|Other|Unknown)\\b', gender_raw, re.IGNORECASE)\n",
1296
+ " fields['gender'] = gender_match.group(1).capitalize() if gender_match else gender_raw\n",
1297
+ " elif line.startswith('4.'):\n",
1298
+ " # Clean and validate profession field\n",
1299
+ " profession_raw = line[2:].strip()\n",
1300
+ " \n",
1301
+ " # Split by comma and validate each profession\n",
1302
+ " professions = [p.strip().lower() for p in profession_raw.split(',')]\n",
1303
+ " valid_profs = []\n",
1304
+ " \n",
1305
+ " for prof in professions:\n",
1306
+ " # Check if it's a valid profession\n",
1307
+ " if prof in VALID_PROFESSIONS:\n",
1308
+ " valid_profs.append(prof)\n",
1309
+ " # Check for common invalid entries\n",
1310
+ " elif prof in ['unknown', '']:\n",
1311
+ " continue\n",
1312
+ " # Reject descriptive words that aren't professions\n",
1313
+ " elif prof in ['sexy', 'photorealistic', 'celebrity', 'famous', 'popular', \n",
1314
+ " 'beautiful', 'attractive', 'hot', 'gorgeous']:\n",
1315
+ " continue\n",
1316
+ " # If it looks like it might be close to a valid profession, keep it\n",
1317
+ " elif any(valid in prof for valid in VALID_PROFESSIONS):\n",
1318
+ " # Try to extract the valid part\n",
1319
+ " for valid in VALID_PROFESSIONS:\n",
1320
+ " if valid in prof:\n",
1321
+ " valid_profs.append(valid)\n",
1322
+ " break\n",
1323
+ " \n",
1324
+ " # Set the cleaned professions or Unknown if none are valid\n",
1325
+ " if valid_profs:\n",
1326
+ " fields['profession_llm'] = ', '.join(valid_profs)\n",
1327
+ " else:\n",
1328
+ " fields['profession_llm'] = 'Unknown'\n",
1329
+ " \n",
1330
+ " elif line.startswith('5.'):\n",
1331
+ " # Clean country field - remove any labels\n",
1332
+ " country_raw = line[2:].strip()\n",
1333
+ " # Remove common prefixes like \"Primary country:\", \"Country:\", etc.\n",
1334
+ " country_raw = re.sub(r'^(Primary\\s+)?(associated\\s+)?country:\\s*', '', country_raw, flags=re.IGNORECASE)\n",
1335
+ " fields['country'] = country_raw\n",
1336
+ " \n",
1337
+ " return fields\n",
1338
+ "\n",
1339
+ "# === PROCESS DATA ===\n",
1340
+ "index_file.parent.mkdir(parents=True, exist_ok=True)\n",
1341
+ "\n",
1342
+ "# Load index\n",
1343
+ "current_index = 0\n",
1344
+ "if index_file.exists():\n",
1345
+ " try:\n",
1346
+ " current_index = int(index_file.read_text().strip())\n",
1347
+ " except:\n",
1348
+ " current_index = 0\n",
1349
+ "\n",
1350
+ "print(f\"Resuming from index {current_index}\")\n",
1351
+ "\n",
1352
+ "start_time = time.time()\n",
1353
+ "\n",
1354
+ "for i in tqdm(range(current_index, len(df)), desc=\"EuroLLM Local\"):\n",
1355
+ "\n",
1356
+ " prompt = df.at[i, \"prompt\"]\n",
1357
+ "\n",
1358
+ " # -------- MODEL QUERY WITH RETRIES --------\n",
1359
+ " response = None\n",
1360
+ " for attempt in range(3):\n",
1361
+ " response = query_eurollm_local(prompt)\n",
1362
+ " \n",
1363
+ " # DEBUG: Print first few responses to see what's happening\n",
1364
+ " if i < 5:\n",
1365
+ " print(f\"\\n=== DEBUG Row {i}, Attempt {attempt+1} ===\")\n",
1366
+ " print(f\"Response length: {len(response) if response else 0}\")\n",
1367
+ " print(f\"Response: {response[:500] if response else 'None'}\")\n",
1368
+ " print(\"=\" * 50)\n",
1369
+ " \n",
1370
+ " # Valid response?\n",
1371
+ " if response and len(response.strip()) > 10:\n",
1372
+ " break\n",
1373
+ " \n",
1374
+ " print(f\"⚠️ Row {i}: Empty or invalid response, retry {attempt+1}/3\")\n",
1375
+ " time.sleep(0.5)\n",
1376
+ "\n",
1377
+ " # If still invalid → DO NOT overwrite previous data\n",
1378
+ " if not response or len(response.strip()) <= 10:\n",
1379
+ " print(f\"❌ Row {i}: failed after retries, not writing, not advancing index\")\n",
1380
+ " continue\n",
1381
+ "\n",
1382
+ " parsed = parse_response(response)\n",
1383
+ "\n",
1384
+ " # DEBUG: Print first few parsed results\n",
1385
+ " if i < 5:\n",
1386
+ " print(f\"\\n=== PARSED Row {i} ===\")\n",
1387
+ " for key, value in parsed.items():\n",
1388
+ " print(f\" {key}: {value}\")\n",
1389
+ " print(\"=\" * 50)\n",
1390
+ "\n",
1391
+ " # Additional safety: skip rows that parsed as all 'Unknown'\n",
1392
+ " if all(v == \"Unknown\" for v in parsed.values()):\n",
1393
+ " print(f\"❌ Row {i}: parsed as all Unknown (likely model crash); skipping.\")\n",
1394
+ " continue\n",
1395
+ "\n",
1396
+ " # -------- WRITE PARSED FIELDS SAFELY --------\n",
1397
+ " for key, value in parsed.items():\n",
1398
+ " df.at[i, key] = value\n",
1399
+ "\n",
1400
+ " # Advance progress ONLY after successful write\n",
1401
+ " current_index = i + 1\n",
1402
+ "\n",
1403
+ " # -------- GPU MEMORY CLEANUP --------\n",
1404
+ " if torch.cuda.is_available():\n",
1405
+ " torch.cuda.empty_cache()\n",
1406
+ " torch.cuda.synchronize()\n",
1407
+ "\n",
1408
+ " # -------- SAVE LIKE YOUR DEEPSEEK VERSION --------\n",
1409
+ " if (i + 1) % SAVE_INTERVAL == 0 or (i + 1) == len(df):\n",
1410
+ " df.to_csv(output_file, index=False)\n",
1411
+ " with open(index_file, \"w\") as f:\n",
1412
+ " f.write(str(current_index))\n",
1413
+ " print(f\"💾 Progress saved after row {i+1}\")\n",
1414
+ "\n",
1415
+ "# Final save\n",
1416
+ "df.to_csv(output_file, index=False)\n",
1417
+ "index_file.write_text(str(current_index))\n",
1418
+ "print(\"✅ Finished full dataset.\")"
1419
+ ]
1420
+ },
1421
+ {
1422
+ "cell_type": "code",
1423
+ "execution_count": null,
1424
+ "id": "a55a5e30-83f3-4f7c-a537-b1216d4e8a07",
1425
+ "metadata": {},
1426
+ "outputs": [],
1427
+ "source": []
1428
+ }
1429
+ ],
1430
+ "metadata": {
1431
+ "kernelspec": {
1432
+ "display_name": "pm-paper",
1433
+ "language": "python",
1434
+ "name": "pm-paper"
1435
+ },
1436
+ "language_info": {
1437
+ "codemirror_mode": {
1438
+ "name": "ipython",
1439
+ "version": 3
1440
+ },
1441
+ "file_extension": ".py",
1442
+ "mimetype": "text/x-python",
1443
+ "name": "python",
1444
+ "nbconvert_exporter": "python",
1445
+ "pygments_lexer": "ipython3",
1446
+ "version": "3.11.13"
1447
+ }
1448
+ },
1449
+ "nbformat": 4,
1450
+ "nbformat_minor": 5
1451
+ }
jupyter_notebooks/.ipynb_checkpoints/Section_2-3-4_Figure_8_Step_2_response_comparison_and_consensus_extraction-checkpoint.ipynb ADDED
The diff for this file is too large to render. See raw diff
 
jupyter_notebooks/.ipynb_checkpoints/Section_2-4_Figure_9_ectract_LoRA_metadata_v2-checkpoint.ipynb ADDED
@@ -0,0 +1,400 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "cells": [
3
+ {
4
+ "cell_type": "markdown",
5
+ "id": "f36422c8",
6
+ "metadata": {},
7
+ "source": [
8
+ "# LoRA metadata"
9
+ ]
10
+ },
11
+ {
12
+ "cell_type": "raw",
13
+ "id": "8a2feb6e",
14
+ "metadata": {},
15
+ "source": [
16
+ "LoRA Metadata Processing Workflow\n",
17
+ "┌──────────────────────┐ ┌───────────────────────────────────────────────────┐\n",
18
+ "│ Load CSV File │ --> │ Read adapter metadata CSV file. │\n",
19
+ "│ Read Model Versions │ │ Extract model version IDs and relevant data. │\n",
20
+ "└─────────┬────────────┘ └───────────────────────────────────────────────────┘\n",
21
+ " ↓\n",
22
+ "┌──────────────────────┐ ┌───────────────────────────────────────────────────┐\n",
23
+ "│ Download Adapter │ --> │ Use stored download URLs to fetch adapter files │\n",
24
+ "│ Files Using API │ │ using rotating API keys. │\n",
25
+ "└─────────┬────────────┘ └───────────────────────────────────────────────────┘\n",
26
+ "┌──────────────────────┐ ┌───────────────────────────────────────────────────┐\n",
27
+ "│ Parse Metadata │ --> │ Extract safetensors metadata, such as training │\n",
28
+ "│ from SafeTensor │ │ images, model type, and architecture. │\n",
29
+ "│ Files │ │ │\n",
30
+ "└─────────┬────────────┘ └───────────────────────────────────────────────────┘\n",
31
+ " ↓\n",
32
+ "┌──────────────────────┐ ┌───────────────────────────────────────────────────┐\n",
33
+ "│ Store Parsed │ --> │ Save extracted metadata into structured JSON │\n",
34
+ "│ Metadata as JSON │ │ files for later analysis. │\n",
35
+ "└─────────┬────────────┘ └───────────────────────────────────────────────────┘\n",
36
+ "\n",
37
+ "┌──────────────────────┐ ┌───────────────────────────────────────────────────┐\n",
38
+ "│ Process JSON Files │ --> │ Read saved JSON metadata, extract relevant │\n",
39
+ "│ for Consolidation │ │ details, and filter necessary attributes. │\n",
40
+ "└─────────┬────────────┘ └───────────────────────────────────────────────────┘\n",
41
+ " ↓\n",
42
+ "┌──────────────────────┐ ┌───────────────────────────────────────────────────┐\n",
43
+ "│ Extract Training │ --> │ Identify most frequent training tags, architectures│\n",
44
+ "│ Tags & Model Info │ │ and systems used for model creation. │\n",
45
+ "└─────────┬────────────┘ └───────────────────────────────────────────────────┘\n",
46
+ " ↓\n",
47
+ "┌──────────────────────┐ ┌───────────────────────────────────────────────────┐\n",
48
+ "│ Save Consolidated │ --> │ Store all processed metadata in a structured CSV │\n",
49
+ "│ Metadata to CSV │ │ format for final analysis. ���\n",
50
+ "└──────────────────────┘ └───────────────────────────────────────────────────┘\n"
51
+ ]
52
+ },
53
+ {
54
+ "cell_type": "code",
55
+ "execution_count": null,
56
+ "id": "efc9939d",
57
+ "metadata": {},
58
+ "outputs": [],
59
+ "source": [
60
+ "import os\n",
61
+ "import re\n",
62
+ "import json\n",
63
+ "import csv\n",
64
+ "import struct\n",
65
+ "import requests\n",
66
+ "from pathlib import Path\n",
67
+ "import pandas as pd\n",
68
+ "from collections import Counter\n",
69
+ "from concurrent.futures import ProcessPoolExecutor\n",
70
+ "from pathlib import Path\n",
71
+ "import matplotlib.pyplot as plt\n",
72
+ "from matplotlib.font_manager import FontProperties\n",
73
+ "from matplotlib import font_manager\n",
74
+ "import pandas as pd\n",
75
+ "from collections import Counter\n",
76
+ "from concurrent.futures import ProcessPoolExecutor\n",
77
+ "\n",
78
+ "# Define the current directory and important file paths\n",
79
+ "current_dir = Path.cwd()\n",
80
+ "\n",
81
+ "# Define frequently used directories\n",
82
+ "\n",
83
+ "data_dir = current_dir.parent / 'data/csv/adapters.csv'\n",
84
+ "fonts_dir = current_dir.parent / 'misc/assets/fonts'\n",
85
+ "plots_dir = current_dir.parent / 'results/plots'\n",
86
+ "raw_data_dir = current_dir.parent / 'data/adapter_metadata/lora' ### location of the LoRA metadata (JSON)\n",
87
+ "temp_dir = current_dir.parent / 'data/raw/adapters_safetensors'\n",
88
+ "misc_dir = current_dir.parent / 'misc'\n",
89
+ "\n",
90
+ "# File paths\n",
91
+ "adapters_csv = current_dir.parent / 'data/csv/adapters.csv'\n",
92
+ "output_json_dir = raw_data_dir\n",
93
+ "api_keys_file = misc_dir / 'credentials/civit.txt'\n",
94
+ "\n",
95
+ "# Ensure directories exist\n",
96
+ "os.makedirs(output_json_dir, exist_ok=True)\n",
97
+ "os.makedirs(temp_dir, exist_ok=True)\n",
98
+ "\n",
99
+ "\n",
100
+ "# Load fonts into Matplotlib\n",
101
+ "for font_path in font_paths:\n",
102
+ " font_manager.fontManager.addfont(font_path)\n",
103
+ "\n",
104
+ "# Set default font family for plots\n",
105
+ "plt.rcParams['font.family'] = ['Noto Sans JP', 'Noto Sans SC', 'sans-serif']\n",
106
+ "\n",
107
+ "print('Paths and fonts initialized successfully.')\n",
108
+ "\n",
109
+ "print('Paths initialized successfully.')"
110
+ ]
111
+ },
112
+ {
113
+ "cell_type": "markdown",
114
+ "id": "87a58593",
115
+ "metadata": {},
116
+ "source": [
117
+ "## Step 2: Download LoRA and extract *.safetensors metadata\n",
118
+ "This script downloads LoRA adapters from the filtered Civiverse-Models dataset and extracts the metadata found within the *.safetensors' data structure"
119
+ ]
120
+ },
121
+ {
122
+ "cell_type": "code",
123
+ "execution_count": null,
124
+ "id": "abd3a0bc",
125
+ "metadata": {},
126
+ "outputs": [],
127
+ "source": [
128
+ "import os\n",
129
+ "import sys\n",
130
+ "import csv\n",
131
+ "import json\n",
132
+ "import struct\n",
133
+ "import time\n",
134
+ "import requests\n",
135
+ "import signal\n",
136
+ "import contextlib\n",
137
+ "from pathlib import Path\n",
138
+ "import re\n",
139
+ "\n",
140
+ "# === Paste your API keys here ===\n",
141
+ "API_KEYS = [\n",
142
+ " \"REDACTED-CIVITAI-API-KEY-1\", #DISCORD\n",
143
+ " \"REDACTED-CIVITAI-API-KEY-2\", #ASDD 1\n",
144
+ " \"REDACTED-CIVITAI-API-KEY-3\", #ASDD 2\n",
145
+ " \"REDACTED-CIVITAI-API-KEY-4\", #BSDD \n",
146
+ " \"REDACTED-CIVITAI-API-KEY-5\"\n",
147
+ "]\n",
148
+ "if not API_KEYS or any(not isinstance(k, str) or not k.strip() for k in API_KEYS):\n",
149
+ " raise ValueError(\"Please paste at least one valid API key into API_KEYS.\")\n",
150
+ "\n",
151
+ "# === Config (adjust paths as needed) ===\n",
152
+ "current_dir = Path.cwd()\n",
153
+ "output_json_dir = current_dir.parent / \"data/adapter_metadata/lora\" # where JSON outputs go\n",
154
+ "temp_dir = current_dir.parent / \"data/raw/adapters_safetensors\" # where downloads go\n",
155
+ "csv_path = current_dir.parent / \"data/csv/adapters_poi_false_sfw.csv\"\n",
156
+ "\n",
157
+ "os.makedirs(output_json_dir, exist_ok=True)\n",
158
+ "os.makedirs(temp_dir, exist_ok=True)\n",
159
+ "\n",
160
+ "# === API key state ===\n",
161
+ "current_key_index = 0\n",
162
+ "\n",
163
+ "\n",
164
+ "\n",
165
+ "def safe_filename(name: str, max_length: int = 100) -> str:\n",
166
+ " # Replace unsafe chars\n",
167
+ " sanitized = re.sub(r'[^a-zA-Z0-9_\\-]', '_', name)\n",
168
+ " # Truncate if too long\n",
169
+ " if len(sanitized) > max_length:\n",
170
+ " sanitized = sanitized[:max_length]\n",
171
+ " return sanitized\n",
172
+ "\n",
173
+ "\n",
174
+ "def get_headers():\n",
175
+ " global current_key_index\n",
176
+ " return {\n",
177
+ " \"Accept\": \"application/json\",\n",
178
+ " \"Authorization\": f\"Bearer {API_KEYS[current_key_index].strip()}\"\n",
179
+ " }\n",
180
+ "\n",
181
+ "def rotate_api_key():\n",
182
+ " global current_key_index\n",
183
+ " if current_key_index < len(API_KEYS) - 1:\n",
184
+ " current_key_index += 1\n",
185
+ " print(f\"🔁 Rotated to API key #{current_key_index + 1}\")\n",
186
+ " else:\n",
187
+ " raise Exception(\"All API keys have been exhausted.\")\n",
188
+ "\n",
189
+ "# === Utilities ===\n",
190
+ "def save_json(data, filename):\n",
191
+ " with open(filename, 'w', encoding=\"utf-8\") as f:\n",
192
+ " json.dump(data, f, indent=4, ensure_ascii=False)\n",
193
+ "\n",
194
+ "def parse_safetensors(file_path):\n",
195
+ " # Minimal, tolerant metadata reader; returns {} on failure.\n",
196
+ " try:\n",
197
+ " with open(file_path, 'rb') as f:\n",
198
+ " file_data = f.read()\n",
199
+ " # Many safetensors use 8-byte header length; this code follows your original logic\n",
200
+ " # (4-byte) but keeps the 8-byte skip. Keep if it's working in your dataset.\n",
201
+ " metadata_size = struct.unpack('<I', file_data[:4])[0]\n",
202
+ " metadata_bytes = file_data[8:8 + metadata_size]\n",
203
+ " metadata_str = metadata_bytes.decode('utf-8', errors='replace')\n",
204
+ " metadata = json.loads(metadata_str)\n",
205
+ " return metadata.get('__metadata__', {})\n",
206
+ " except Exception as e:\n",
207
+ " print(f\"Error parsing safetensors file: {e}\")\n",
208
+ " return {}\n",
209
+ "\n",
210
+ "# === Timeout context ===\n",
211
+ "class TimeoutException(Exception):\n",
212
+ " pass\n",
213
+ "\n",
214
+ "@contextlib.contextmanager\n",
215
+ "def time_limit(seconds):\n",
216
+ " def signal_handler(signum, frame):\n",
217
+ " raise TimeoutException(f\"Timed out after {seconds} seconds\")\n",
218
+ " # Note: SIGALRM works on Unix-like OS; on Windows this will be a no-op.\n",
219
+ " try:\n",
220
+ " signal.signal(signal.SIGALRM, signal_handler)\n",
221
+ " signal.alarm(seconds)\n",
222
+ " except Exception:\n",
223
+ " # Fallback: no hard alarm on non-Unix systems\n",
224
+ " pass\n",
225
+ " try:\n",
226
+ " yield\n",
227
+ " finally:\n",
228
+ " try:\n",
229
+ " signal.alarm(0)\n",
230
+ " except Exception:\n",
231
+ " pass\n",
232
+ "\n",
233
+ "# === Download with timeout, retries, backoff, and key rotation ===\n",
234
+ "def download_file(url, output_folder, timeout=30, overall_timeout=120, max_retries=3):\n",
235
+ " filename = url.split(\"/\")[-1]\n",
236
+ " output_path = os.path.join(output_folder, filename)\n",
237
+ "\n",
238
+ " global current_key_index\n",
239
+ " retries = 0\n",
240
+ " backoff = 2\n",
241
+ "\n",
242
+ " while current_key_index < len(API_KEYS):\n",
243
+ " try:\n",
244
+ " with time_limit(overall_timeout): # global cap per download\n",
245
+ " #print(f\"➡️ GET {url} using key #{current_key_index + 1}\")\n",
246
+ " resp = requests.get(\n",
247
+ " url,\n",
248
+ " headers=get_headers(),\n",
249
+ " stream=True,\n",
250
+ " timeout=(10, timeout), # (connect timeout, per-chunk read timeout)\n",
251
+ " )\n",
252
+ "\n",
253
+ " # Auth errors → rotate key\n",
254
+ " if resp.status_code in (401, 403):\n",
255
+ " print(f\"❌ Auth {resp.status_code} with key #{current_key_index + 1}. Rotating.\")\n",
256
+ " rotate_api_key()\n",
257
+ " retries = 0\n",
258
+ " backoff = 2\n",
259
+ " continue\n",
260
+ "\n",
261
+ " # Not found → bubble up as FileNotFoundError (do not rotate)\n",
262
+ " if resp.status_code == 404:\n",
263
+ " raise FileNotFoundError(f\"Model not found at {url}\")\n",
264
+ "\n",
265
+ " # Rate limit → either rotate or wait/backoff\n",
266
+ " if resp.status_code == 429:\n",
267
+ " print(\"⏳ Rate limited (429).\", end=\" \")\n",
268
+ " if current_key_index < len(API_KEYS) - 1:\n",
269
+ " print(\"Rotating key.\")\n",
270
+ " rotate_api_key()\n",
271
+ " retries = 0\n",
272
+ " backoff = 2\n",
273
+ " continue\n",
274
+ " else:\n",
275
+ " print(f\"Waiting {backoff}s (no other keys).\")\n",
276
+ " time.sleep(backoff)\n",
277
+ " backoff = min(backoff * 2, 60)\n",
278
+ " continue\n",
279
+ "\n",
280
+ " # Other HTTP errors → raise to RequestException path\n",
281
+ " resp.raise_for_status()\n",
282
+ "\n",
283
+ " # Save file\n",
284
+ " with open(output_path, 'wb') as fh:\n",
285
+ " for chunk in resp.iter_content(chunk_size=8192):\n",
286
+ " if chunk:\n",
287
+ " fh.write(chunk)\n",
288
+ "\n",
289
+ " return output_path, filename\n",
290
+ "\n",
291
+ " except TimeoutException as e:\n",
292
+ " # Hard overall timeout → propagate\n",
293
+ " raise e\n",
294
+ " except requests.exceptions.RequestException as e:\n",
295
+ " # Network-ish errors: retry same key with backoff up to max_retries\n",
296
+ " retries += 1\n",
297
+ " if retries <= max_retries:\n",
298
+ " print(f\"🌐 Network error (try {retries}/{max_retries}) with key #{current_key_index + 1}: {e}\")\n",
299
+ " time.sleep(backoff)\n",
300
+ " backoff = min(backoff * 2, 60)\n",
301
+ " continue\n",
302
+ " else:\n",
303
+ " raise Exception(f\"Failed to download {url} after {max_retries} retries: {e}\")\n",
304
+ "\n",
305
+ " # If we exit the loop, we truly ran out\n",
306
+ " raise Exception(\"All API keys have been exhausted or failed.\")\n",
307
+ "\n",
308
+ "# === Main processing ===\n",
309
+ "def process_csv(csv_path):\n",
310
+ " with open(csv_path, newline='', encoding='utf-8') as csvfile:\n",
311
+ " reader = csv.DictReader(csvfile)\n",
312
+ " for index, row in enumerate(reader):\n",
313
+ " # Collect up to 20 version IDs; use the most recent\n",
314
+ " version_ids = []\n",
315
+ " for i in range(1, 21):\n",
316
+ " k = f'version_id_{i}'\n",
317
+ " if k in row and row[k]:\n",
318
+ " try:\n",
319
+ " version_ids.append(int(float(row[k])))\n",
320
+ " except ValueError:\n",
321
+ " print(f\"Invalid version_id value '{row[k]}' in row: {row}\")\n",
322
+ "\n",
323
+ " if not version_ids:\n",
324
+ " print(f\"No valid version IDs found in row: {row}\")\n",
325
+ " continue\n",
326
+ "\n",
327
+ " most_recent_version_id = str(max(version_ids))\n",
328
+ " name = row.get('name', 'unknown')\n",
329
+ " sanitized_name = safe_filename(name, max_length=100)\n",
330
+ " new_json_file = os.path.join(\n",
331
+ " output_json_dir,\n",
332
+ " f\"{index:08d}_{most_recent_version_id}_{sanitized_name}.json\"\n",
333
+ " )\n",
334
+ "\n",
335
+ " # Skip if JSON already exists\n",
336
+ " if os.path.exists(new_json_file):\n",
337
+ " #print(f\"↩️ Skipping versionID {most_recent_version_id} (JSON already exists)\")\n",
338
+ " continue\n",
339
+ "\n",
340
+ " try:\n",
341
+ " adapter_file, fname = download_file(\n",
342
+ " row['downloadUrl'], str(temp_dir),\n",
343
+ " timeout=30, overall_timeout=180\n",
344
+ " )\n",
345
+ " metadata = parse_safetensors(adapter_file)\n",
346
+ "\n",
347
+ " civitaidata = {\n",
348
+ " k: (int(v) if str(v).isdigit() else v)\n",
349
+ " for k, v in row.items()\n",
350
+ " }\n",
351
+ " new_json_data = {\n",
352
+ " \"civitaidata\": civitaidata,\n",
353
+ " \"metadata\": metadata,\n",
354
+ " \"versionID\": most_recent_version_id\n",
355
+ " }\n",
356
+ " save_json(new_json_data, new_json_file)\n",
357
+ " #print(f\"✅ Created JSON for versionID {most_recent_version_id} with file {fname}\")\n",
358
+ "\n",
359
+ " except FileNotFoundError as e:\n",
360
+ " print(f\"⚠️ {e} — saving empty metadata.\")\n",
361
+ " civitaidata = {\n",
362
+ " k: (int(v) if str(v).isdigit() else v)\n",
363
+ " for k, v in row.items()\n",
364
+ " }\n",
365
+ " empty_json = {\n",
366
+ " \"civitaidata\": civitaidata,\n",
367
+ " \"metadata\": {},\n",
368
+ " \"versionID\": most_recent_version_id,\n",
369
+ " \"error\": \"Model not found (404)\"\n",
370
+ " }\n",
371
+ " save_json(empty_json, new_json_file)\n",
372
+ " except Exception as e:\n",
373
+ " print(f\"⚠️ Error processing versionID {most_recent_version_id}: {e}\")\n",
374
+ " civitaidata = {\n",
375
+ " k: (int(v) if str(v).isdigit() else v)\n",
376
+ " for k, v in row.items()\n",
377
+ " }\n",
378
+ " empty_json = {\n",
379
+ " \"civitaidata\": civitaidata,\n",
380
+ " \"metadata\": {},\n",
381
+ " \"versionID\": most_recent_version_id,\n",
382
+ " \"error\": str(e)\n",
383
+ " }\n",
384
+ " save_json(empty_json, new_json_file)\n",
385
+ " print(f\"💾 Saved empty JSON for versionID {most_recent_version_id} due to failure.\")\n",
386
+ "\n",
387
+ "# === Run ===\n",
388
+ "if __name__ == \"__main__\":\n",
389
+ " process_csv(csv_path)\n"
390
+ ]
391
+ }
392
+ ],
393
+ "metadata": {
394
+ "language_info": {
395
+ "name": "python"
396
+ }
397
+ },
398
+ "nbformat": 4,
399
+ "nbformat_minor": 5
400
+ }
jupyter_notebooks/Section_2-3-4_Figure_8_Step_1_LLM_annotation.ipynb CHANGED
@@ -35,71 +35,10 @@
35
  },
36
  {
37
  "cell_type": "code",
38
- "execution_count": 3,
39
  "id": "a287eef4",
40
  "metadata": {},
41
- "outputs": [
42
- {
43
- "name": "stdout",
44
- "output_type": "stream",
45
- "text": [
46
- "✅ spaCy model loaded: en_core_web_sm\n",
47
- "Loaded 50861 rows\n",
48
- "\n",
49
- "🔄 Processing names with spaCy NER...\n",
50
- "\n",
51
- "📊 Name cleaning examples (with spaCy NER):\n",
52
- "====================================================================================================\n",
53
- "Original Name | Cleaned Name \n",
54
- "====================================================================================================\n",
55
- "Super Pose Book Vol.1 - ControlNet | Super Pose Book \n",
56
- "Liyuu LoRA | Liyuu \n",
57
- "HashimotoKanna/ 橋本環奈 _JP_Actress | HashimotoKanna \n",
58
- "Emma Watson (JG) | Emma Watson \n",
59
- "Gal Gadot「LoRa」 | Gal Gadot \n",
60
- "Scarlett Johansson「LoRa」 | Scarlett Johansson \n",
61
- "Gakki | Aragaki Yui | 新垣結衣 | Gakki \n",
62
- "Actress Satomi_石原○○ | Actress Satomi \n",
63
- "Game of Thrones Cast | Game Thrones Cast \n",
64
- "Natalie Portman「LoRa」 | Natalie Portman \n",
65
- "Emma Watson LoRA | Emma Watson \n",
66
- "Karina Makina Lora | Karina Makina \n",
67
- "WRAV YUA_三xx亜 | WRAV YUA \n",
68
- "Dilraba Dilmurat 迪丽热巴 | Dilraba Dilmurat \n",
69
- "MIMI,大幂幂 | MIMI \n",
70
- "Chinese Idol - YangMi杨幂 | Chinese Idol YangMi杨幂 \n",
71
- "Xiaorouseeu / 小柔SeeU - Chinese cosplayer and influencer | Xiaorouseeu \n",
72
- "Jennifer Connelly (80s/90s) | Jennifer Connelly \n",
73
- "====================================================================================================\n",
74
- "\n",
75
- "🧪 Leetspeak translation examples:\n",
76
- " 4kira LoRA -> akira\n",
77
- " 3mma Watson v2 -> Watson\n",
78
- " 1rene LORA -> irene\n",
79
- " L3vi Ackerman -> Levi Ackerman\n",
80
- "\n",
81
- "📈 Statistics:\n",
82
- " Total rows: 50861\n",
83
- " Non-empty names: 50858\n",
84
- " Empty names: 3\n",
85
- "\n",
86
- "🎯 Sample spaCy NER results:\n",
87
- " 1. IU\n",
88
- " 2. Super Pose Book\n",
89
- " 3. Liyuu\n",
90
- " 4. Irene\n",
91
- " 5. AESPA Karina\n",
92
- " 6. Saika Kawakita\n",
93
- " 7. Liu Yifei\n",
94
- " 8. HashimotoKanna\n",
95
- " 9. Emma Watson\n",
96
- " 10. Gal Gadot\n",
97
- "\n",
98
- "✅ Cleaned 50861 names using spaCy NER\n",
99
- "💾 Saved to /home/lauhp/000_PHD/000_010_PUBLICATION/CODE/pm-paper/data/CSV/model_adapter/real_person_adapter_step_01_NER.csv\n"
100
- ]
101
- }
102
- ],
103
  "source": [
104
  "import pandas as pd\n",
105
  "import re\n",
@@ -1011,13 +950,928 @@
1011
  "# Test mode (first 100 rows)\n",
1012
  "# annotate_dataset(model_type='mistral', test_mode=True, test_size=100)\n"
1013
  ]
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1014
  }
1015
  ],
1016
  "metadata": {
1017
  "kernelspec": {
1018
- "display_name": "latm",
1019
  "language": "python",
1020
- "name": "python3"
1021
  },
1022
  "language_info": {
1023
  "codemirror_mode": {
@@ -1029,7 +1883,7 @@
1029
  "name": "python",
1030
  "nbconvert_exporter": "python",
1031
  "pygments_lexer": "ipython3",
1032
- "version": "3.10.15"
1033
  }
1034
  },
1035
  "nbformat": 4,
 
35
  },
36
  {
37
  "cell_type": "code",
38
+ "execution_count": null,
39
  "id": "a287eef4",
40
  "metadata": {},
41
+ "outputs": [],
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
42
  "source": [
43
  "import pandas as pd\n",
44
  "import re\n",
 
950
  "# Test mode (first 100 rows)\n",
951
  "# annotate_dataset(model_type='mistral', test_mode=True, test_size=100)\n"
952
  ]
953
+ },
954
+ {
955
+ "cell_type": "markdown",
956
+ "id": "6431d347-d80c-4e8b-83a7-531e5df95a72",
957
+ "metadata": {},
958
+ "source": [
959
+ "## EuroLLM-9B-Instruct"
960
+ ]
961
+ },
962
+ {
963
+ "cell_type": "code",
964
+ "execution_count": null,
965
+ "id": "e8203abc-e7c3-4cb6-aaeb-fdc6933981fc",
966
+ "metadata": {},
967
+ "outputs": [],
968
+ "source": [
969
+ "import pandas as pd\n",
970
+ "import json\n",
971
+ "import time\n",
972
+ "import re\n",
973
+ "from pathlib import Path\n",
974
+ "from tqdm import tqdm\n",
975
+ "import torch\n",
976
+ "from transformers import AutoModelForCausalLM, AutoTokenizer, BitsAndBytesConfig\n",
977
+ "import signal\n",
978
+ "from contextlib import contextmanager\n",
979
+ "\n",
980
+ "current_dir = Path.cwd()\n",
981
+ "input_file = current_dir.parent / \"data/CSV/model_adapter/real_person_adapter_step_02_NER.csv\"\n",
982
+ "professions_file = current_dir.parent / \"misc/lists/professions.csv\"\n",
983
+ "professions_mapped_file = current_dir.parent / \"misc/lists/professions_mapped.csv\"\n",
984
+ "# === PROCESS DATA ===\n",
985
+ "\n",
986
+ "\n",
987
+ "# === CONFIGURATION ===\n",
988
+ "TEST_MODE = False\n",
989
+ "TEST_SIZE = 100\n",
990
+ "MAX_ROWS = 50862\n",
991
+ "SAVE_INTERVAL = 10\n",
992
+ "\n",
993
+ "\n",
994
+ "index_file = current_dir.parent / \"misc/query_indicies/eurollm_local_query_index.txt\"\n",
995
+ "output_file = current_dir.parent / f\"data/CSV/eurollm_local_annotated_POI{'_test' if TEST_MODE else ''}.csv\"\n",
996
+ "\n",
997
+ "# Model settings\n",
998
+ "MODEL_NAME = \"utter-project/EuroLLM-9B-Instruct\"\n",
999
+ "#MODEL_NAME = \"Qwen/Qwen2.5-32B-Instruct\"\n",
1000
+ "#MODEL_NAME = \"Qwen/Qwen2.5-14B-Instruct\"\n",
1001
+ "#MODEL_NAME = \"Qwen/Qwen3-235B-A22B-Instruct-2507-FP8\"\n",
1002
+ "#MODEL_NAME = \"mistralai/Mistral-7B-Instruct-v0.3\"\n",
1003
+ "CACHE_DIR = current_dir.parent / \"data/models\"\n",
1004
+ "CACHE_DIR.mkdir(parents=True, exist_ok=True)\n",
1005
+ "\n",
1006
+ "# Define the SPECIFIC profession categories\n",
1007
+ "PROFESSION_CATEGORIES = [\n",
1008
+ " \"actor\",\n",
1009
+ " \"adult performer\",\n",
1010
+ " \"singer/musician\",\n",
1011
+ " \"model\",\n",
1012
+ " \"online personality\",\n",
1013
+ " \"public figure\",\n",
1014
+ " \"voice actor/ASMR\",\n",
1015
+ " \"sports professional\",\n",
1016
+ " \"tv personality\"\n",
1017
+ "]\n",
1018
+ "\n",
1019
+ "# === LOAD MODEL ===\n",
1020
+ "print(f\"Loading model: {MODEL_NAME}\")\n",
1021
+ "print(f\"Cache directory: {CACHE_DIR}\")\n",
1022
+ "print(f\"This may take a while on first run...\\n\")\n",
1023
+ "\n",
1024
+ "# Check GPU availability\n",
1025
+ "device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n",
1026
+ "print(f\"Device: {device}\")\n",
1027
+ "\n",
1028
+ "if device == \"cpu\":\n",
1029
+ " print(\"⚠️ WARNING: No GPU detected! Inference will be VERY slow.\")\n",
1030
+ " print(\" Consider using a GPU or reducing model size.\")\n",
1031
+ "\n",
1032
+ "# Get HF token from credentials file\n",
1033
+ "import os\n",
1034
+ "credentials_dir = current_dir.parent / \"misc/credentials\"\n",
1035
+ "hf_token_file = credentials_dir / \"hf_token.txt\"\n",
1036
+ "\n",
1037
+ "HF_TOKEN = None\n",
1038
+ "if hf_token_file.exists():\n",
1039
+ " HF_TOKEN = hf_token_file.read_text().strip()\n",
1040
+ " print(\"✅ HF token loaded from credentials file\")\n",
1041
+ "else:\n",
1042
+ " print(\"⚠️ HF token file not found at:\", hf_token_file)\n",
1043
+ " print(\" The script will try to use cached credentials from 'huggingface-cli login'\")\n",
1044
+ " print(\" Or create the file: misc/credentials/hf_token.txt with your token\")\n",
1045
+ " HF_TOKEN = None # Will use cached token if available\n",
1046
+ "\n",
1047
+ "# Load tokenizer\n",
1048
+ "print(\"Loading tokenizer...\")\n",
1049
+ "try:\n",
1050
+ " tokenizer = AutoTokenizer.from_pretrained(\n",
1051
+ " MODEL_NAME,\n",
1052
+ " cache_dir=str(CACHE_DIR),\n",
1053
+ " use_fast=True,\n",
1054
+ " token=HF_TOKEN\n",
1055
+ " )\n",
1056
+ "except Exception as e:\n",
1057
+ " print(f\"Failed with use_fast=True, trying use_fast=False...\")\n",
1058
+ " tokenizer = AutoTokenizer.from_pretrained(\n",
1059
+ " MODEL_NAME,\n",
1060
+ " cache_dir=str(CACHE_DIR),\n",
1061
+ " use_fast=False,\n",
1062
+ " token=HF_TOKEN\n",
1063
+ " )\n",
1064
+ "\n",
1065
+ "# Ensure pad token is set\n",
1066
+ "if tokenizer.pad_token is None:\n",
1067
+ " tokenizer.pad_token = tokenizer.eos_token\n",
1068
+ "\n",
1069
+ "print(\"✅ Tokenizer loaded\")\n",
1070
+ "\n",
1071
+ "# Configure 8-bit quantization for A100\n",
1072
+ "print(\"Configuring 8-bit quantization...\")\n",
1073
+ "quantization_config = BitsAndBytesConfig(\n",
1074
+ " load_in_8bit=True,\n",
1075
+ " llm_int8_threshold=6.0,\n",
1076
+ " llm_int8_has_fp16_weight=False\n",
1077
+ ")\n",
1078
+ "\n",
1079
+ "# Load model with 8-bit quantization\n",
1080
+ "print(\"Loading model with 8-bit quantization (this may take several minutes)...\")\n",
1081
+ "model = AutoModelForCausalLM.from_pretrained(\n",
1082
+ " MODEL_NAME,\n",
1083
+ " cache_dir=str(CACHE_DIR),\n",
1084
+ " quantization_config=quantization_config,\n",
1085
+ " device_map=\"auto\",\n",
1086
+ " trust_remote_code=False,\n",
1087
+ " token=HF_TOKEN\n",
1088
+ ")\n",
1089
+ "model.eval()\n",
1090
+ "print(\"✅ Model loaded with 8-bit quantization\")\n",
1091
+ "\n",
1092
+ "# Check VRAM usage\n",
1093
+ "if torch.cuda.is_available():\n",
1094
+ " vram_gb = torch.cuda.max_memory_allocated() / 1024**3\n",
1095
+ " print(f\"VRAM used: {vram_gb:.2f} GB\\n\")\n",
1096
+ "\n",
1097
+ "# === LOAD DATA ===\n",
1098
+ "if output_file.exists():\n",
1099
+ " print(\"Loading annotated CSV...\")\n",
1100
+ " df = pd.read_csv(output_file)\n",
1101
+ "else:\n",
1102
+ " print(\"Loading raw input CSV...\")\n",
1103
+ " df = pd.read_csv(input_file)\n",
1104
+ "\n",
1105
+ "\n",
1106
+ "# Try to load profession mapping files\n",
1107
+ "try:\n",
1108
+ " professions_df = pd.read_csv(professions_file)\n",
1109
+ " print(f\"✅ Loaded professions.csv\")\n",
1110
+ "except:\n",
1111
+ " print(\"⚠️ Warning: professions.csv not found\")\n",
1112
+ "\n",
1113
+ "try:\n",
1114
+ " prof_mapped_df = pd.read_csv(professions_mapped_file)\n",
1115
+ " print(f\"✅ Loaded profession mapping with {len(prof_mapped_df)} categories\")\n",
1116
+ "except:\n",
1117
+ " print(\"⚠️ Warning: professions_mapped.csv not found, using default categories\")\n",
1118
+ "\n",
1119
+ "profession_str = \", \".join(PROFESSION_CATEGORIES)\n",
1120
+ "\n",
1121
+ "print(f\"Loaded {len(df)} rows\")\n",
1122
+ "print(f\"\\nProfession categories ({len(PROFESSION_CATEGORIES)}):\")\n",
1123
+ "for cat in PROFESSION_CATEGORIES:\n",
1124
+ " print(f\" - {cat}\")\n",
1125
+ "\n",
1126
+ "if TEST_MODE:\n",
1127
+ " print(f\"\\nRunning in TEST MODE with {TEST_SIZE} samples\")\n",
1128
+ " df = df.head(TEST_SIZE).copy()\n",
1129
+ "elif MAX_ROWS:\n",
1130
+ " df = df.head(MAX_ROWS).copy()\n",
1131
+ "\n",
1132
+ "# === CREATE PROMPTS (OPTIMIZED FOR CLEAN OUTPUTS) ===\n",
1133
+ "def create_prompt(row):\n",
1134
+ " \"\"\"Create prompt for EuroLLM annotation with strict formatting requirements.\"\"\"\n",
1135
+ " name = row['real_name'] if pd.notna(row.get('real_name')) else row.get('name', '')\n",
1136
+ " \n",
1137
+ " # Gather hints\n",
1138
+ " hints = []\n",
1139
+ " if pd.notna(row.get('likely_profession')):\n",
1140
+ " hints.append(str(row['likely_profession']))\n",
1141
+ " if pd.notna(row.get('likely_nationality')):\n",
1142
+ " hints.append(str(row['likely_nationality']))\n",
1143
+ " if pd.notna(row.get('likely_country')):\n",
1144
+ " hints.append(str(row['likely_country']))\n",
1145
+ " \n",
1146
+ " # Add tags if we don't have enough hints\n",
1147
+ " if len(hints) < 3:\n",
1148
+ " for i in range(1, 8):\n",
1149
+ " tag_col = f'tag_{i}'\n",
1150
+ " if tag_col in row and pd.notna(row[tag_col]):\n",
1151
+ " tag_val = str(row[tag_col])\n",
1152
+ " if tag_val not in hints:\n",
1153
+ " hints.append(tag_val)\n",
1154
+ " if len(hints) >= 5:\n",
1155
+ " break\n",
1156
+ " \n",
1157
+ " hint_text = \", \".join(hints[:5]) if hints else \"none\"\n",
1158
+ " \n",
1159
+ " return f\"\"\"Extract information about '{name}'. \n",
1160
+ "Context hints (DO NOT copy these as professions): {hint_text}\n",
1161
+ "\n",
1162
+ "Respond with EXACTLY 5 numbered lines. Each line must contain ONLY the value, no labels or extra text.\n",
1163
+ "\n",
1164
+ "FORMAT REQUIREMENTS:\n",
1165
+ "1. Full legal name in Western order (first last). VALUE ONLY.\n",
1166
+ "2. Stage names/aliases, comma-separated. If none, write \"None\". VALUE ONLY.\n",
1167
+ "3. Gender: MUST be exactly one word: Male, Female, Other, or Unknown. VALUE ONLY.\n",
1168
+ "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",
1169
+ "5. Primary country: Country name only (e.g., \"China\", \"United States\", \"Colombia\"). VALUE ONLY.\n",
1170
+ "\n",
1171
+ "CRITICAL RULES FOR PROFESSIONS (Line 4):\n",
1172
+ "- ONLY use the exact profession categories listed above\n",
1173
+ "- DO NOT use descriptive words like \"sexy\", \"photorealistic\", \"celebrity\"\n",
1174
+ "- DO NOT copy the hint words as professions\n",
1175
+ "- If uncertain about profession, write \"Unknown\"\n",
1176
+ "- Valid professions are ONLY: actor, adult performer, singer/musician, model, online personality, public figure, voice actor/ASMR, sports professional, tv personality\n",
1177
+ "- Actress = actor, streamer = online personality, YouTuber = online personality\n",
1178
+ "\n",
1179
+ "OTHER RULES:\n",
1180
+ "- Use \"Unknown\" when uncertain or for fictional characters\n",
1181
+ "- NO explanatory text, NO labels like \"Gender:\", NO prefixes\n",
1182
+ "- For multi-role people, list up to 3 categories by relevance\"\"\"\n",
1183
+ "\n",
1184
+ "# Create prompts\n",
1185
+ "print(\"\\nCreating prompts...\")\n",
1186
+ "df['prompt'] = df.apply(create_prompt, axis=1)\n",
1187
+ "print(\"✅ Prompts created\")\n",
1188
+ "\n",
1189
+ "@contextmanager\n",
1190
+ "def timeout(duration):\n",
1191
+ " def handler(signum, frame):\n",
1192
+ " raise TimeoutError(\"Operation timed out\")\n",
1193
+ " \n",
1194
+ " # Set the signal handler and alarm\n",
1195
+ " signal.signal(signal.SIGALRM, handler)\n",
1196
+ " signal.alarm(duration)\n",
1197
+ " try:\n",
1198
+ " yield\n",
1199
+ " finally:\n",
1200
+ " signal.alarm(0) # Disable the alarm\n",
1201
+ "\n",
1202
+ "\n",
1203
+ "def query_eurollm_local(prompt: str) -> str:\n",
1204
+ " \"\"\"Query EuroLLM locally via transformers with very low temperature.\"\"\"\n",
1205
+ " try:\n",
1206
+ " # Format as chat message for EuroLLM with strict system prompt\n",
1207
+ " messages = [\n",
1208
+ " {\"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",
1209
+ " {\"role\": \"user\", \"content\": prompt}\n",
1210
+ " ]\n",
1211
+ " \n",
1212
+ " # Tokenize\n",
1213
+ " if hasattr(tokenizer, 'apply_chat_template') and tokenizer.chat_template is not None:\n",
1214
+ " text = tokenizer.apply_chat_template(\n",
1215
+ " messages,\n",
1216
+ " tokenize=False,\n",
1217
+ " add_generation_prompt=True\n",
1218
+ " )\n",
1219
+ " else:\n",
1220
+ " # Fallback for models without chat template\n",
1221
+ " text = f\"[INST] {prompt} [/INST]\"\n",
1222
+ " \n",
1223
+ " inputs = tokenizer([text], return_tensors=\"pt\", padding=True).to(device)\n",
1224
+ " \n",
1225
+ " # Generate with timeout and very low temperature\n",
1226
+ " with timeout(60):\n",
1227
+ " with torch.no_grad():\n",
1228
+ " outputs = model.generate(\n",
1229
+ " **inputs,\n",
1230
+ " max_new_tokens=100,\n",
1231
+ " temperature=0.01, # Very low temperature for more deterministic outputs\n",
1232
+ " do_sample=True, # Must be True when temperature is set\n",
1233
+ " pad_token_id=tokenizer.eos_token_id\n",
1234
+ " )\n",
1235
+ " \n",
1236
+ " # Decode\n",
1237
+ " generated_ids = outputs[0][inputs['input_ids'].shape[1]:]\n",
1238
+ " response = tokenizer.decode(generated_ids, skip_special_tokens=True)\n",
1239
+ " \n",
1240
+ " return response.strip()\n",
1241
+ " \n",
1242
+ " except TimeoutError:\n",
1243
+ " print(f\"[ERROR] Generation timed out after 60 seconds\")\n",
1244
+ " return None\n",
1245
+ " except Exception as e:\n",
1246
+ " print(f\"Generation error: {e}\")\n",
1247
+ " import traceback\n",
1248
+ " traceback.print_exc()\n",
1249
+ " return None\n",
1250
+ "\n",
1251
+ " \n",
1252
+ "# === PARSE RESPONSE WITH CLEANING ===\n",
1253
+ "def parse_response(response):\n",
1254
+ " \"\"\"Parse EuroLLM response into structured fields with cleaning.\"\"\"\n",
1255
+ " if not response:\n",
1256
+ " return {\n",
1257
+ " 'full_name': 'Unknown',\n",
1258
+ " 'aliases': 'Unknown',\n",
1259
+ " 'gender': 'Unknown',\n",
1260
+ " 'profession_llm': 'Unknown',\n",
1261
+ " 'country': 'Unknown'\n",
1262
+ " }\n",
1263
+ " \n",
1264
+ " # Valid profession categories\n",
1265
+ " VALID_PROFESSIONS = {\n",
1266
+ " \"actor\", \"adult performer\", \"singer/musician\", \"model\", \n",
1267
+ " \"online personality\", \"public figure\", \"voice actor/asmr\", \n",
1268
+ " \"sports professional\", \"tv personality\"\n",
1269
+ " }\n",
1270
+ " \n",
1271
+ " # Split into lines and clean\n",
1272
+ " lines = [line.strip() for line in response.split('\\n') if line.strip()]\n",
1273
+ " \n",
1274
+ " # Initialize with Unknown values\n",
1275
+ " fields = {\n",
1276
+ " 'full_name': 'Unknown',\n",
1277
+ " 'aliases': 'Unknown',\n",
1278
+ " 'gender': 'Unknown',\n",
1279
+ " 'profession_llm': 'Unknown',\n",
1280
+ " 'country': 'Unknown'\n",
1281
+ " }\n",
1282
+ " \n",
1283
+ " # Extract information from each numbered line\n",
1284
+ " for line in lines:\n",
1285
+ " if line.startswith('1.'):\n",
1286
+ " fields['full_name'] = line[2:].strip()\n",
1287
+ " elif line.startswith('2.'):\n",
1288
+ " fields['aliases'] = line[2:].strip()\n",
1289
+ " elif line.startswith('3.'):\n",
1290
+ " # Clean gender field - remove any labels\n",
1291
+ " gender_raw = line[2:].strip()\n",
1292
+ " # Remove common prefixes\n",
1293
+ " gender_raw = re.sub(r'^(Gender:|gender:)\\s*', '', gender_raw, flags=re.IGNORECASE)\n",
1294
+ " # Extract just the gender word\n",
1295
+ " gender_match = re.search(r'\\b(Male|Female|Other|Unknown)\\b', gender_raw, re.IGNORECASE)\n",
1296
+ " fields['gender'] = gender_match.group(1).capitalize() if gender_match else gender_raw\n",
1297
+ " elif line.startswith('4.'):\n",
1298
+ " # Clean and validate profession field\n",
1299
+ " profession_raw = line[2:].strip()\n",
1300
+ " \n",
1301
+ " # Split by comma and validate each profession\n",
1302
+ " professions = [p.strip().lower() for p in profession_raw.split(',')]\n",
1303
+ " valid_profs = []\n",
1304
+ " \n",
1305
+ " for prof in professions:\n",
1306
+ " # Check if it's a valid profession\n",
1307
+ " if prof in VALID_PROFESSIONS:\n",
1308
+ " valid_profs.append(prof)\n",
1309
+ " # Check for common invalid entries\n",
1310
+ " elif prof in ['unknown', '']:\n",
1311
+ " continue\n",
1312
+ " # Reject descriptive words that aren't professions\n",
1313
+ " elif prof in ['sexy', 'photorealistic', 'celebrity', 'famous', 'popular', \n",
1314
+ " 'beautiful', 'attractive', 'hot', 'gorgeous']:\n",
1315
+ " continue\n",
1316
+ " # If it looks like it might be close to a valid profession, keep it\n",
1317
+ " elif any(valid in prof for valid in VALID_PROFESSIONS):\n",
1318
+ " # Try to extract the valid part\n",
1319
+ " for valid in VALID_PROFESSIONS:\n",
1320
+ " if valid in prof:\n",
1321
+ " valid_profs.append(valid)\n",
1322
+ " break\n",
1323
+ " \n",
1324
+ " # Set the cleaned professions or Unknown if none are valid\n",
1325
+ " if valid_profs:\n",
1326
+ " fields['profession_llm'] = ', '.join(valid_profs)\n",
1327
+ " else:\n",
1328
+ " fields['profession_llm'] = 'Unknown'\n",
1329
+ " \n",
1330
+ " elif line.startswith('5.'):\n",
1331
+ " # Clean country field - remove any labels\n",
1332
+ " country_raw = line[2:].strip()\n",
1333
+ " # Remove common prefixes like \"Primary country:\", \"Country:\", etc.\n",
1334
+ " country_raw = re.sub(r'^(Primary\\s+)?(associated\\s+)?country:\\s*', '', country_raw, flags=re.IGNORECASE)\n",
1335
+ " fields['country'] = country_raw\n",
1336
+ " \n",
1337
+ " return fields\n",
1338
+ "\n",
1339
+ "# === PROCESS DATA ===\n",
1340
+ "index_file.parent.mkdir(parents=True, exist_ok=True)\n",
1341
+ "\n",
1342
+ "# Load index\n",
1343
+ "current_index = 0\n",
1344
+ "if index_file.exists():\n",
1345
+ " try:\n",
1346
+ " current_index = int(index_file.read_text().strip())\n",
1347
+ " except:\n",
1348
+ " current_index = 0\n",
1349
+ "\n",
1350
+ "print(f\"Resuming from index {current_index}\")\n",
1351
+ "\n",
1352
+ "start_time = time.time()\n",
1353
+ "\n",
1354
+ "for i in tqdm(range(current_index, len(df)), desc=\"EuroLLM Local\"):\n",
1355
+ "\n",
1356
+ " prompt = df.at[i, \"prompt\"]\n",
1357
+ "\n",
1358
+ " # -------- MODEL QUERY WITH RETRIES --------\n",
1359
+ " response = None\n",
1360
+ " for attempt in range(3):\n",
1361
+ " response = query_eurollm_local(prompt)\n",
1362
+ " \n",
1363
+ " # DEBUG: Print first few responses to see what's happening\n",
1364
+ " if i < 5:\n",
1365
+ " print(f\"\\n=== DEBUG Row {i}, Attempt {attempt+1} ===\")\n",
1366
+ " print(f\"Response length: {len(response) if response else 0}\")\n",
1367
+ " print(f\"Response: {response[:500] if response else 'None'}\")\n",
1368
+ " print(\"=\" * 50)\n",
1369
+ " \n",
1370
+ " # Valid response?\n",
1371
+ " if response and len(response.strip()) > 10:\n",
1372
+ " break\n",
1373
+ " \n",
1374
+ " print(f\"⚠️ Row {i}: Empty or invalid response, retry {attempt+1}/3\")\n",
1375
+ " time.sleep(0.5)\n",
1376
+ "\n",
1377
+ " # If still invalid → DO NOT overwrite previous data\n",
1378
+ " if not response or len(response.strip()) <= 10:\n",
1379
+ " print(f\"❌ Row {i}: failed after retries, not writing, not advancing index\")\n",
1380
+ " continue\n",
1381
+ "\n",
1382
+ " parsed = parse_response(response)\n",
1383
+ "\n",
1384
+ " # DEBUG: Print first few parsed results\n",
1385
+ " if i < 5:\n",
1386
+ " print(f\"\\n=== PARSED Row {i} ===\")\n",
1387
+ " for key, value in parsed.items():\n",
1388
+ " print(f\" {key}: {value}\")\n",
1389
+ " print(\"=\" * 50)\n",
1390
+ "\n",
1391
+ " # Additional safety: skip rows that parsed as all 'Unknown'\n",
1392
+ " if all(v == \"Unknown\" for v in parsed.values()):\n",
1393
+ " print(f\"❌ Row {i}: parsed as all Unknown (likely model crash); skipping.\")\n",
1394
+ " continue\n",
1395
+ "\n",
1396
+ " # -------- WRITE PARSED FIELDS SAFELY --------\n",
1397
+ " for key, value in parsed.items():\n",
1398
+ " df.at[i, key] = value\n",
1399
+ "\n",
1400
+ " # Advance progress ONLY after successful write\n",
1401
+ " current_index = i + 1\n",
1402
+ "\n",
1403
+ " # -------- GPU MEMORY CLEANUP --------\n",
1404
+ " if torch.cuda.is_available():\n",
1405
+ " torch.cuda.empty_cache()\n",
1406
+ " torch.cuda.synchronize()\n",
1407
+ "\n",
1408
+ " # -------- SAVE LIKE YOUR DEEPSEEK VERSION --------\n",
1409
+ " if (i + 1) % SAVE_INTERVAL == 0 or (i + 1) == len(df):\n",
1410
+ " df.to_csv(output_file, index=False)\n",
1411
+ " with open(index_file, \"w\") as f:\n",
1412
+ " f.write(str(current_index))\n",
1413
+ " print(f\"💾 Progress saved after row {i+1}\")\n",
1414
+ "\n",
1415
+ "# Final save\n",
1416
+ "df.to_csv(output_file, index=False)\n",
1417
+ "index_file.write_text(str(current_index))\n",
1418
+ "print(\"✅ Finished full dataset.\")"
1419
+ ]
1420
+ },
1421
+ {
1422
+ "cell_type": "markdown",
1423
+ "id": "472e5ac2-ec04-4bfa-8a67-116277238c15",
1424
+ "metadata": {},
1425
+ "source": [
1426
+ "## Mistral 24b instruct"
1427
+ ]
1428
+ },
1429
+ {
1430
+ "cell_type": "code",
1431
+ "execution_count": 1,
1432
+ "id": "a55a5e30-83f3-4f7c-a537-b1216d4e8a07",
1433
+ "metadata": {
1434
+ "execution": {
1435
+ "iopub.execute_input": "2025-12-08T23:57:35.685431Z",
1436
+ "iopub.status.busy": "2025-12-08T23:57:35.685314Z",
1437
+ "iopub.status.idle": "2025-12-08T23:59:48.656498Z",
1438
+ "shell.execute_reply": "2025-12-08T23:59:48.655927Z",
1439
+ "shell.execute_reply.started": "2025-12-08T23:57:35.685419Z"
1440
+ }
1441
+ },
1442
+ "outputs": [
1443
+ {
1444
+ "name": "stderr",
1445
+ "output_type": "stream",
1446
+ "text": [
1447
+ "/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",
1448
+ " from .autonotebook import tqdm as notebook_tqdm\n"
1449
+ ]
1450
+ },
1451
+ {
1452
+ "name": "stdout",
1453
+ "output_type": "stream",
1454
+ "text": [
1455
+ "Loading model: mistralai/Mistral-Small-Instruct-2409\n",
1456
+ "Cache directory: /shares/weddigen.ki.uzh/laura_wagner/phase_01/pm-paper/data/models\n",
1457
+ "This may take a while on first run (~65GB download)...\n",
1458
+ "\n",
1459
+ "Device: cuda\n",
1460
+ "Loading tokenizer...\n",
1461
+ "✅ Tokenizer loaded\n"
1462
+ ]
1463
+ },
1464
+ {
1465
+ "ename": "NameError",
1466
+ "evalue": "name 'BitsAndBytesConfig' is not defined",
1467
+ "output_type": "error",
1468
+ "traceback": [
1469
+ "\u001b[31m---------------------------------------------------------------------------\u001b[39m",
1470
+ "\u001b[31mNameError\u001b[39m Traceback (most recent call last)",
1471
+ "\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",
1472
+ "\u001b[31mNameError\u001b[39m: name 'BitsAndBytesConfig' is not defined"
1473
+ ]
1474
+ }
1475
+ ],
1476
+ "source": [
1477
+ "import pandas as pd\n",
1478
+ "import json\n",
1479
+ "import time\n",
1480
+ "import re\n",
1481
+ "from pathlib import Path\n",
1482
+ "from tqdm import tqdm\n",
1483
+ "import torch\n",
1484
+ "from transformers import AutoModelForCausalLM, AutoTokenizer\n",
1485
+ "\n",
1486
+ "current_dir = Path.cwd()\n",
1487
+ "input_file = current_dir.parent / \"data/CSV/model_adapter/real_person_adapter_step_02_NER.csv\"\n",
1488
+ "professions_file = current_dir.parent / \"misc/lists/professions.csv\"\n",
1489
+ "professions_mapped_file = current_dir.parent / \"misc/lists/professions_mapped.csv\"\n",
1490
+ "# === PROCESS DATA ===\n",
1491
+ "\n",
1492
+ "\n",
1493
+ "# === CONFIGURATION ===\n",
1494
+ "TEST_MODE = False\n",
1495
+ "TEST_SIZE = 100\n",
1496
+ "MAX_ROWS = 50862\n",
1497
+ "SAVE_INTERVAL = 10\n",
1498
+ "\n",
1499
+ "output_file = current_dir.parent / f\"data/CSV/mistral24_local_annotated_POI{'_test' if TEST_MODE else ''}.csv\"\n",
1500
+ "index_file = current_dir.parent / \"misc/query_indicies/mistral24_local_query_index.txt\"\n",
1501
+ "\n",
1502
+ "\n",
1503
+ "# Model settings\n",
1504
+ "#MODEL_NAME = \"mistralai/Mistral-Small-3.1-24B-Instruct-2503\"\n",
1505
+ "MODEL_NAME = \"mistralai/Mistral-Small-Instruct-2409\"\n",
1506
+ "#MODEL_NAME = \"mistralai/Mistral-7B-Instruct-v0.3\"\n",
1507
+ "CACHE_DIR = current_dir.parent / \"data/models\"\n",
1508
+ "CACHE_DIR.mkdir(parents=True, exist_ok=True)\n",
1509
+ "\n",
1510
+ "# Define the SPECIFIC profession categories\n",
1511
+ "PROFESSION_CATEGORIES = [\n",
1512
+ " \"actor\",\n",
1513
+ " \"adult performer\",\n",
1514
+ " \"singer/musician\",\n",
1515
+ " \"model\",\n",
1516
+ " \"online personality\",\n",
1517
+ " \"public figure\",\n",
1518
+ " \"voice actor/ASMR\",\n",
1519
+ " \"sports professional\",\n",
1520
+ " \"tv personality\"\n",
1521
+ "]\n",
1522
+ "\n",
1523
+ "# === LOAD MODEL ===\n",
1524
+ "print(f\"Loading model: {MODEL_NAME}\")\n",
1525
+ "print(f\"Cache directory: {CACHE_DIR}\")\n",
1526
+ "print(f\"This may take a while on first run (~65GB download)...\\n\")\n",
1527
+ "\n",
1528
+ "# Check GPU availability\n",
1529
+ "device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n",
1530
+ "print(f\"Device: {device}\")\n",
1531
+ "\n",
1532
+ "if device == \"cpu\":\n",
1533
+ " print(\"⚠️ WARNING: No GPU detected! Inference will be VERY slow.\")\n",
1534
+ " print(\" Consider using a GPU or reducing model size.\")\n",
1535
+ "\n",
1536
+ "# Load tokenizer\n",
1537
+ "print(\"Loading tokenizer...\")\n",
1538
+ "try:\n",
1539
+ " tokenizer = AutoTokenizer.from_pretrained(\n",
1540
+ " MODEL_NAME,\n",
1541
+ " cache_dir=str(CACHE_DIR),\n",
1542
+ " use_fast=True\n",
1543
+ " )\n",
1544
+ "except Exception as e:\n",
1545
+ " print(f\"Failed with use_fast=True, trying use_fast=False...\")\n",
1546
+ " tokenizer = AutoTokenizer.from_pretrained(\n",
1547
+ " MODEL_NAME,\n",
1548
+ " cache_dir=str(CACHE_DIR),\n",
1549
+ " use_fast=False\n",
1550
+ " )\n",
1551
+ "\n",
1552
+ "# Ensure pad token is set\n",
1553
+ "if tokenizer.pad_token is None:\n",
1554
+ " tokenizer.pad_token = tokenizer.eos_token\n",
1555
+ "\n",
1556
+ "print(\"✅ Tokenizer loaded\")\n",
1557
+ "\n",
1558
+ "quantization_config = BitsAndBytesConfig(\n",
1559
+ " load_in_8bit=True\n",
1560
+ ")\n",
1561
+ "\n",
1562
+ "\n",
1563
+ "# Load model with optimizations\n",
1564
+ "print(\"Loading model (this may take several minutes)...\")\n",
1565
+ "model = AutoModelForCausalLM.from_pretrained(\n",
1566
+ " MODEL_NAME,\n",
1567
+ " cache_dir=str(CACHE_DIR),\n",
1568
+ " torch_dtype=torch.bfloat16,\n",
1569
+ " quantization_config=quantization_config,\n",
1570
+ " device_map=\"auto\",\n",
1571
+ " trust_remote_code=False\n",
1572
+ ")\n",
1573
+ "model.eval()\n",
1574
+ "print(\"✅ Model loaded\")\n",
1575
+ "\n",
1576
+ "# Check VRAM usage\n",
1577
+ "if torch.cuda.is_available():\n",
1578
+ " vram_gb = torch.cuda.max_memory_allocated() / 1024**3\n",
1579
+ " print(f\"VRAM used: {vram_gb:.2f} GB\\n\")\n",
1580
+ "\n",
1581
+ "# === LOAD DATA ===\n",
1582
+ "print(\"Loading raw input CSV...\")\n",
1583
+ "df = pd.read_csv(input_file) # ALWAYS load the full input\n",
1584
+ "print(f\"Loaded {len(df)} rows from input file\")\n",
1585
+ "\n",
1586
+ "# If we have previous annotations, merge them\n",
1587
+ "if output_file.exists():\n",
1588
+ " print(\"Found existing annotations, merging...\")\n",
1589
+ " existing_df = pd.read_csv(output_file)\n",
1590
+ " print(f\"Existing annotations has {len(existing_df)} rows\")\n",
1591
+ " \n",
1592
+ " # Update df with existing annotations\n",
1593
+ " # Only update the columns that were annotated\n",
1594
+ " annotation_cols = ['full_name', 'aliases', 'gender', 'profession_llm', 'country']\n",
1595
+ " for col in annotation_cols:\n",
1596
+ " if col in existing_df.columns:\n",
1597
+ " df[col] = existing_df[col][:len(df)] # Make sure we don't exceed df length\n",
1598
+ " \n",
1599
+ " print(f\"Merged annotations, continuing with {len(df)} total rows\")\n",
1600
+ "\n",
1601
+ "\n",
1602
+ "# Try to load profession mapping files\n",
1603
+ "try:\n",
1604
+ " professions_df = pd.read_csv(professions_file)\n",
1605
+ " print(f\"✅ Loaded professions.csv\")\n",
1606
+ "except:\n",
1607
+ " print(\"⚠️ Warning: professions.csv not found\")\n",
1608
+ "\n",
1609
+ "try:\n",
1610
+ " prof_mapped_df = pd.read_csv(professions_mapped_file)\n",
1611
+ " print(f\"✅ Loaded profession mapping with {len(prof_mapped_df)} categories\")\n",
1612
+ "except:\n",
1613
+ " print(\"⚠️ Warning: professions_mapped.csv not found, using default categories\")\n",
1614
+ "\n",
1615
+ "profession_str = \", \".join(PROFESSION_CATEGORIES)\n",
1616
+ "\n",
1617
+ "print(f\"Loaded {len(df)} rows\")\n",
1618
+ "print(f\"\\nProfession categories ({len(PROFESSION_CATEGORIES)}):\")\n",
1619
+ "for cat in PROFESSION_CATEGORIES:\n",
1620
+ " print(f\" - {cat}\")\n",
1621
+ "\n",
1622
+ "if TEST_MODE:\n",
1623
+ " print(f\"\\nRunning in TEST MODE with {TEST_SIZE} samples\")\n",
1624
+ " df = df.head(TEST_SIZE).copy()\n",
1625
+ "elif MAX_ROWS:\n",
1626
+ " df = df.head(MAX_ROWS).copy()\n",
1627
+ "\n",
1628
+ "# === CREATE PROMPTS (DEEPSEEK STYLE) ===\n",
1629
+ "def create_prompt(row):\n",
1630
+ " \"\"\"Create prompt for Mistral annotation with specific profession categories.\"\"\"\n",
1631
+ " name = row['real_name'] if pd.notna(row.get('real_name')) else row.get('name', '')\n",
1632
+ " \n",
1633
+ " # Gather hints\n",
1634
+ " hints = []\n",
1635
+ " if pd.notna(row.get('likely_profession')):\n",
1636
+ " hints.append(str(row['likely_profession']))\n",
1637
+ " if pd.notna(row.get('likely_nationality')):\n",
1638
+ " hints.append(str(row['likely_nationality']))\n",
1639
+ " if pd.notna(row.get('likely_country')):\n",
1640
+ " hints.append(str(row['likely_country']))\n",
1641
+ " \n",
1642
+ " # Add tags if we don't have enough hints\n",
1643
+ " if len(hints) < 3:\n",
1644
+ " for i in range(1, 8):\n",
1645
+ " tag_col = f'tag_{i}'\n",
1646
+ " if tag_col in row and pd.notna(row[tag_col]):\n",
1647
+ " tag_val = str(row[tag_col])\n",
1648
+ " if tag_val not in hints:\n",
1649
+ " hints.append(tag_val)\n",
1650
+ " if len(hints) >= 5:\n",
1651
+ " break\n",
1652
+ " \n",
1653
+ " hint_text = \", \".join(hints[:5]) if hints else \"none\"\n",
1654
+ " \n",
1655
+ " return f\"\"\"Given '{name}' ({hint_text}), provide:\n",
1656
+ "1. Full legal name (Western order if non-latin script)\n",
1657
+ "2. Any stage names/aliases (comma separated)\n",
1658
+ "3. Gender (Male/Female/Other/Unknown)\n",
1659
+ "4. Top 3 most likely professions from ONLY these categories:\n",
1660
+ " - actor\n",
1661
+ " - adult performer\n",
1662
+ " - singer/musician\n",
1663
+ " - model\n",
1664
+ " - online personality (includes streamers, cosplayers, influencers)\n",
1665
+ " - public figure (includes politicians, activists, journalists, authors)\n",
1666
+ " - voice actor/ASMR\n",
1667
+ " - sports professional\n",
1668
+ " - tv personality (includes hosts, presenters, reality TV)\n",
1669
+ "\n",
1670
+ "5. Primary country associated\n",
1671
+ "\n",
1672
+ "IMPORTANT:\n",
1673
+ "- Choose professions ONLY from the 9 categories above\n",
1674
+ "- Provide up to 3 professions, comma-separated, ordered by relevance\n",
1675
+ "- Be SPECIFIC: choose the most accurate category for each role\n",
1676
+ "- \"online personality\" includes: streamers, cosplayers, YouTubers, influencers, content creators\n",
1677
+ "- Use 'Unknown' when uncertain or for fictional characters/places\n",
1678
+ "- For multi-role people, list all relevant categories (e.g., \"actor, singer/musician, online personality\")\n",
1679
+ "- For country respond with one word only, for example China or Columbia\n",
1680
+ "- actress = actor\n",
1681
+ "\n",
1682
+ "Respond with exactly 5 numbered lines.\"\"\"\n",
1683
+ "\n",
1684
+ "# Create prompts\n",
1685
+ "print(\"\\nCreating prompts...\")\n",
1686
+ "df['prompt'] = df.apply(create_prompt, axis=1)\n",
1687
+ "print(\"✅ Prompts created\")\n",
1688
+ "\n",
1689
+ "# === QUERY MISTRAL LOCAL ===\n",
1690
+ "def query_mistral_local(prompt: str) -> str:\n",
1691
+ " \"\"\"Query Mistral locally via transformers.\"\"\"\n",
1692
+ " try:\n",
1693
+ " # Format as chat message for Mistral\n",
1694
+ " messages = [\n",
1695
+ " {\"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",
1696
+ " {\"role\": \"user\", \"content\": prompt}\n",
1697
+ " ]\n",
1698
+ " \n",
1699
+ " # Tokenize\n",
1700
+ " if hasattr(tokenizer, 'apply_chat_template'):\n",
1701
+ " text = tokenizer.apply_chat_template(\n",
1702
+ " messages,\n",
1703
+ " tokenize=False,\n",
1704
+ " add_generation_prompt=True\n",
1705
+ " )\n",
1706
+ " else:\n",
1707
+ " # Fallback for older tokenizers\n",
1708
+ " text = f\"[INST] {prompt} [/INST]\"\n",
1709
+ " \n",
1710
+ " inputs = tokenizer([text], return_tensors=\"pt\", padding=True).to(device)\n",
1711
+ " \n",
1712
+ " # Generate\n",
1713
+ " with torch.no_grad():\n",
1714
+ " outputs = model.generate(\n",
1715
+ " **inputs,\n",
1716
+ " max_new_tokens=512,\n",
1717
+ " temperature=0.05,\n",
1718
+ " do_sample=True,\n",
1719
+ " top_p=0.8,\n",
1720
+ " pad_token_id=tokenizer.pad_token_id if tokenizer.pad_token_id else tokenizer.eos_token_id\n",
1721
+ " )\n",
1722
+ " \n",
1723
+ " # Decode\n",
1724
+ " generated_ids = outputs[0][inputs['input_ids'].shape[1]:]\n",
1725
+ " response = tokenizer.decode(generated_ids, skip_special_tokens=True)\n",
1726
+ " \n",
1727
+ " return response.strip()\n",
1728
+ " \n",
1729
+ " except Exception as e:\n",
1730
+ " print(f\"Generation error: {e}\")\n",
1731
+ " return None\n",
1732
+ "\n",
1733
+ "# === PARSE RESPONSE (DEEPSEEK STYLE) ===\n",
1734
+ "def parse_response(response):\n",
1735
+ " \"\"\"Parse Mistral response into structured fields.\"\"\"\n",
1736
+ " if not response:\n",
1737
+ " return {\n",
1738
+ " 'full_name': 'Unknown',\n",
1739
+ " 'aliases': 'Unknown',\n",
1740
+ " 'gender': 'Unknown',\n",
1741
+ " 'profession_llm': 'Unknown',\n",
1742
+ " 'country': 'Unknown'\n",
1743
+ " }\n",
1744
+ " \n",
1745
+ " # Split into lines and clean\n",
1746
+ " lines = [line.strip() for line in response.split('\\n') if line.strip()]\n",
1747
+ " \n",
1748
+ " # Initialize with Unknown values\n",
1749
+ " fields = {\n",
1750
+ " 'full_name': 'Unknown',\n",
1751
+ " 'aliases': 'Unknown',\n",
1752
+ " 'gender': 'Unknown',\n",
1753
+ " 'profession_llm': 'Unknown',\n",
1754
+ " 'country': 'Unknown'\n",
1755
+ " }\n",
1756
+ " \n",
1757
+ " # Extract information from each numbered line\n",
1758
+ " for line in lines:\n",
1759
+ " if line.startswith('1.'):\n",
1760
+ " fields['full_name'] = line[2:].strip()\n",
1761
+ " elif line.startswith('2.'):\n",
1762
+ " fields['aliases'] = line[2:].strip()\n",
1763
+ " elif line.startswith('3.'):\n",
1764
+ " fields['gender'] = line[2:].strip()\n",
1765
+ " elif line.startswith('4.'):\n",
1766
+ " fields['profession_llm'] = line[2:].strip()\n",
1767
+ " elif line.startswith('5.'):\n",
1768
+ " fields['country'] = line[2:].strip()\n",
1769
+ " \n",
1770
+ " return fields\n",
1771
+ "\n",
1772
+ "# === PROCESS DATA ===\n",
1773
+ "output_file = current_dir.parent / f\"data/CSV/mistral24_local_annotated_POI{'_test' if TEST_MODE else ''}.csv\"\n",
1774
+ "index_file = current_dir.parent / \"misc/query_indicies/mistral24_local_query_index.txt\"\n",
1775
+ "\n",
1776
+ "index_file.parent.mkdir(parents=True, exist_ok=True)\n",
1777
+ "\n",
1778
+ "# Load index\n",
1779
+ "current_index = 0\n",
1780
+ "if index_file.exists():\n",
1781
+ " try:\n",
1782
+ " current_index = int(index_file.read_text().strip())\n",
1783
+ " except:\n",
1784
+ " current_index = 0\n",
1785
+ "\n",
1786
+ "print(f\"Resuming from index {current_index}\")\n",
1787
+ "\n",
1788
+ "start_time = time.time()\n",
1789
+ "\n",
1790
+ "for i in tqdm(range(current_index, len(df)), desc=\"Mistral Local\"):\n",
1791
+ "\n",
1792
+ " prompt = df.at[i, \"prompt\"]\n",
1793
+ "\n",
1794
+ " # -------- MODEL QUERY WITH RETRIES --------\n",
1795
+ " response = None\n",
1796
+ " for attempt in range(3):\n",
1797
+ " response = query_mistral_local(prompt)\n",
1798
+ " \n",
1799
+ " # Valid response?\n",
1800
+ " if response and len(response.strip()) > 10:\n",
1801
+ " break\n",
1802
+ " \n",
1803
+ " print(f\"⚠️ Row {i}: Empty or invalid response, retry {attempt+1}/3\")\n",
1804
+ " time.sleep(0.5)\n",
1805
+ "\n",
1806
+ " # If still invalid → DO NOT overwrite previous data\n",
1807
+ " if not response or len(response.strip()) <= 10:\n",
1808
+ " print(f\"❌ Row {i}: failed after retries, not writing, not advancing index\")\n",
1809
+ " continue\n",
1810
+ "\n",
1811
+ " parsed = parse_response(response)\n",
1812
+ "\n",
1813
+ " # Additional safety: skip rows that parsed as all 'Unknown'\n",
1814
+ " if all(v == \"Unknown\" for v in parsed.values()):\n",
1815
+ " print(f\"❌ Row {i}: parsed as all Unknown (likely model crash); skipping.\")\n",
1816
+ " continue\n",
1817
+ "\n",
1818
+ " # -------- WRITE PARSED FIELDS SAFELY --------\n",
1819
+ " for key, value in parsed.items():\n",
1820
+ " df.at[i, key] = value\n",
1821
+ "\n",
1822
+ " # Advance progress ONLY after successful write\n",
1823
+ " current_index = i + 1\n",
1824
+ "\n",
1825
+ " # -------- GPU MEMORY CLEANUP --------\n",
1826
+ " if torch.cuda.is_available():\n",
1827
+ " torch.cuda.empty_cache()\n",
1828
+ " torch.cuda.synchronize()\n",
1829
+ "\n",
1830
+ " # -------- SAVE LIKE YOUR DEEPSEEK VERSION --------\n",
1831
+ " if (i + 1) % SAVE_INTERVAL == 0 or (i + 1) == len(df):\n",
1832
+ " df.to_csv(output_file, index=False)\n",
1833
+ " with open(index_file, \"w\") as f:\n",
1834
+ " f.write(str(current_index))\n",
1835
+ " print(f\"💾 Progress saved after row {i+1}\")\n",
1836
+ "\n",
1837
+ "# Final save\n",
1838
+ "df.to_csv(output_file, index=False)\n",
1839
+ "index_file.write_text(str(current_index))\n",
1840
+ "print(\"✅ Finished full dataset.\")\n"
1841
+ ]
1842
+ },
1843
+ {
1844
+ "cell_type": "code",
1845
+ "execution_count": null,
1846
+ "id": "d7212e75-0ff6-45a0-8695-c4a3d3e02818",
1847
+ "metadata": {},
1848
+ "outputs": [],
1849
+ "source": [
1850
+ "import transformers\n",
1851
+ "print(f\"Transformers version: {transformers.__version__}\")\n",
1852
+ "\n",
1853
+ "# Check if Mistral3 is available\n",
1854
+ "try:\n",
1855
+ " from transformers import Mistral3ForCausalLM\n",
1856
+ " print(\"✅ Mistral3 is available\")\n",
1857
+ "except ImportError:\n",
1858
+ " print(\"❌ Mistral3 not available in this transformers version\")"
1859
+ ]
1860
+ },
1861
+ {
1862
+ "cell_type": "code",
1863
+ "execution_count": null,
1864
+ "id": "a6ab032e-246e-4c4e-9776-ff0bfbf6fd9c",
1865
+ "metadata": {},
1866
+ "outputs": [],
1867
+ "source": []
1868
  }
1869
  ],
1870
  "metadata": {
1871
  "kernelspec": {
1872
+ "display_name": "pm-paper",
1873
  "language": "python",
1874
+ "name": "pm-paper"
1875
  },
1876
  "language_info": {
1877
  "codemirror_mode": {
 
1883
  "name": "python",
1884
  "nbconvert_exporter": "python",
1885
  "pygments_lexer": "ipython3",
1886
+ "version": "3.11.13"
1887
  }
1888
  },
1889
  "nbformat": 4,
jupyter_notebooks/Section_2-4_Figure_9_ectract_LoRA_metadata_v2.ipynb CHANGED
@@ -391,8 +391,22 @@
391
  }
392
  ],
393
  "metadata": {
 
 
 
 
 
394
  "language_info": {
395
- "name": "python"
 
 
 
 
 
 
 
 
 
396
  }
397
  },
398
  "nbformat": 4,
 
391
  }
392
  ],
393
  "metadata": {
394
+ "kernelspec": {
395
+ "display_name": "Python 3 (ipykernel)",
396
+ "language": "python",
397
+ "name": "python3"
398
+ },
399
  "language_info": {
400
+ "codemirror_mode": {
401
+ "name": "ipython",
402
+ "version": 3
403
+ },
404
+ "file_extension": ".py",
405
+ "mimetype": "text/x-python",
406
+ "name": "python",
407
+ "nbconvert_exporter": "python",
408
+ "pygments_lexer": "ipython3",
409
+ "version": "3.13.9"
410
  }
411
  },
412
  "nbformat": 4,
misc/query_indicies/eurollm_local_query_index.txt ADDED
@@ -0,0 +1 @@
 
 
1
+ 80
misc/query_indicies/mistral24_local_query_index.txt ADDED
@@ -0,0 +1 @@
 
 
1
+ 3890