Laura Wagner commited on
Commit
d317593
·
1 Parent(s): eb286ec

refactored model inference code

Browse files
jupyter_notebooks/GEMMA_3.ipynb DELETED
@@ -1,400 +0,0 @@
1
- {
2
- "cells": [
3
- {
4
- "cell_type": "code",
5
- "execution_count": null,
6
- "id": "471e0cef-678e-4403-8eca-f8e1991d86de",
7
- "metadata": {},
8
- "outputs": [],
9
- "source": [
10
- "import pandas as pd\n",
11
- "import json\n",
12
- "import time\n",
13
- "import re\n",
14
- "from pathlib import Path\n",
15
- "from tqdm import tqdm\n",
16
- "import torch\n",
17
- "from transformers import AutoModelForCausalLM, AutoTokenizer\n",
18
- "\n",
19
- "current_dir = Path.cwd()\n",
20
- "input_file = current_dir.parent / \"data/CSV/model_adapter/real_person_adapter_step_02_NER.csv\"\n",
21
- "professions_file = current_dir.parent / \"misc/lists/professions.csv\"\n",
22
- "professions_mapped_file = current_dir.parent / \"misc/lists/professions_mapped.csv\"\n",
23
- "# === PROCESS DATA ===\n",
24
- "\n",
25
- "\n",
26
- "# === CONFIGURATION ===\n",
27
- "TEST_MODE = False\n",
28
- "TEST_SIZE = 100\n",
29
- "MAX_ROWS = 50862\n",
30
- "SAVE_INTERVAL = 10\n",
31
- "\n",
32
- "output_file = current_dir.parent / f\"data/CSV/gemma_local_annotated_POI{'_test' if TEST_MODE else ''}.csv\"\n",
33
- "index_file = current_dir.parent / \"misc/query_indicies/gemma_local_query_index.txt\"\n",
34
- "\n",
35
- "# Model settings\n",
36
- "MODEL_NAME = MODEL_NAME = \"google/gemma-3-27b-it\"\n",
37
- "#MODEL_NAME = \"mistralai/Mistral-7B-Instruct-v0.3\"\n",
38
- "CACHE_DIR = current_dir.parent / \"data/models\"\n",
39
- "CACHE_DIR.mkdir(parents=True, exist_ok=True)\n",
40
- "\n",
41
- "# Define the SPECIFIC profession categories\n",
42
- "PROFESSION_CATEGORIES = [\n",
43
- " \"actor\",\n",
44
- " \"adult performer\",\n",
45
- " \"singer/musician\",\n",
46
- " \"model\",\n",
47
- " \"online personality\",\n",
48
- " \"public figure\",\n",
49
- " \"voice actor/ASMR\",\n",
50
- " \"sports professional\",\n",
51
- " \"tv personality\"\n",
52
- "]\n",
53
- "\n",
54
- "# === LOAD MODEL ===\n",
55
- "print(f\"Loading model: {MODEL_NAME}\")\n",
56
- "print(f\"Cache directory: {CACHE_DIR}\")\n",
57
- "print(f\"This may take a while on first run (~65GB download)...\\n\")\n",
58
- "\n",
59
- "# Check GPU availability\n",
60
- "device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n",
61
- "print(f\"Device: {device}\")\n",
62
- "\n",
63
- "if device == \"cpu\":\n",
64
- " print(\"⚠️ WARNING: No GPU detected! Inference will be VERY slow.\")\n",
65
- " print(\" Consider using a GPU or reducing model size.\")\n",
66
- "\n",
67
- "# Load tokenizer\n",
68
- "print(\"Loading tokenizer...\")\n",
69
- "try:\n",
70
- " tokenizer = AutoTokenizer.from_pretrained(\n",
71
- " MODEL_NAME,\n",
72
- " cache_dir=str(CACHE_DIR),\n",
73
- " use_fast=True\n",
74
- " )\n",
75
- "except Exception as e:\n",
76
- " print(f\"Failed with use_fast=True, trying use_fast=False...\")\n",
77
- " tokenizer = AutoTokenizer.from_pretrained(\n",
78
- " MODEL_NAME,\n",
79
- " cache_dir=str(CACHE_DIR),\n",
80
- " use_fast=False\n",
81
- " )\n",
82
- "\n",
83
- "# Ensure pad token is set\n",
84
- "if tokenizer.pad_token is None:\n",
85
- " tokenizer.pad_token = tokenizer.eos_token\n",
86
- "\n",
87
- "print(\"✅ Tokenizer loaded\")\n",
88
- "\n",
89
- "# Load model with optimizations\n",
90
- "print(\"Loading model (this may take several minutes)...\")\n",
91
- "model = AutoModelForCausalLM.from_pretrained(\n",
92
- " MODEL_NAME,\n",
93
- " cache_dir=str(CACHE_DIR),\n",
94
- " torch_dtype=torch.bfloat16,\n",
95
- " device_map=\"auto\",\n",
96
- " trust_remote_code=False\n",
97
- ")\n",
98
- "model.eval()\n",
99
- "print(\"✅ Model loaded\")\n",
100
- "\n",
101
- "# Check VRAM usage\n",
102
- "if torch.cuda.is_available():\n",
103
- " vram_gb = torch.cuda.max_memory_allocated() / 1024**3\n",
104
- " print(f\"VRAM used: {vram_gb:.2f} GB\\n\")\n",
105
- "\n",
106
- "# === LOAD DATA ===\n",
107
- "print(\"Loading raw input CSV...\")\n",
108
- "df = pd.read_csv(input_file) # ALWAYS load the full input\n",
109
- "print(f\"Loaded {len(df)} rows from input file\")\n",
110
- "\n",
111
- "# If we have previous annotations, merge them\n",
112
- "if output_file.exists():\n",
113
- " print(\"Found existing annotations, merging...\")\n",
114
- " existing_df = pd.read_csv(output_file)\n",
115
- " print(f\"Existing annotations has {len(existing_df)} rows\")\n",
116
- " \n",
117
- " # Update df with existing annotations\n",
118
- " # Only update the columns that were annotated\n",
119
- " annotation_cols = ['full_name', 'aliases', 'gender', 'profession_llm', 'country']\n",
120
- " for col in annotation_cols:\n",
121
- " if col in existing_df.columns:\n",
122
- " df[col] = existing_df[col][:len(df)] # Make sure we don't exceed df length\n",
123
- " \n",
124
- " print(f\"Merged annotations, continuing with {len(df)} total rows\")\n",
125
- "\n",
126
- "\n",
127
- "# Try to load profession mapping files\n",
128
- "try:\n",
129
- " professions_df = pd.read_csv(professions_file)\n",
130
- " print(f\"✅ Loaded professions.csv\")\n",
131
- "except:\n",
132
- " print(\"⚠️ Warning: professions.csv not found\")\n",
133
- "\n",
134
- "try:\n",
135
- " prof_mapped_df = pd.read_csv(professions_mapped_file)\n",
136
- " print(f\"✅ Loaded profession mapping with {len(prof_mapped_df)} categories\")\n",
137
- "except:\n",
138
- " print(\"⚠️ Warning: professions_mapped.csv not found, using default categories\")\n",
139
- "\n",
140
- "profession_str = \", \".join(PROFESSION_CATEGORIES)\n",
141
- "\n",
142
- "print(f\"Loaded {len(df)} rows\")\n",
143
- "print(f\"\\nProfession categories ({len(PROFESSION_CATEGORIES)}):\")\n",
144
- "for cat in PROFESSION_CATEGORIES:\n",
145
- " print(f\" - {cat}\")\n",
146
- "\n",
147
- "if TEST_MODE:\n",
148
- " print(f\"\\nRunning in TEST MODE with {TEST_SIZE} samples\")\n",
149
- " df = df.head(TEST_SIZE).copy()\n",
150
- "elif MAX_ROWS:\n",
151
- " df = df.head(MAX_ROWS).copy()\n",
152
- "\n",
153
- "# === CREATE PROMPTS ===\n",
154
- "def create_prompt(row):\n",
155
- " \"\"\"Create prompt for Gemma annotation with specific profession categories.\"\"\"\n",
156
- " name = row['real_name'] if pd.notna(row.get('real_name')) else row.get('name', '')\n",
157
- " \n",
158
- " # Gather hints\n",
159
- " hints = []\n",
160
- " if pd.notna(row.get('likely_profession')):\n",
161
- " hints.append(str(row['likely_profession']))\n",
162
- " if pd.notna(row.get('likely_nationality')):\n",
163
- " hints.append(str(row['likely_nationality']))\n",
164
- " if pd.notna(row.get('likely_country')):\n",
165
- " hints.append(str(row['likely_country']))\n",
166
- " \n",
167
- " # Add tags if we don't have enough hints\n",
168
- " if len(hints) < 3:\n",
169
- " for i in range(1, 8):\n",
170
- " tag_col = f'tag_{i}'\n",
171
- " if tag_col in row and pd.notna(row[tag_col]):\n",
172
- " tag_val = str(row[tag_col])\n",
173
- " if tag_val not in hints:\n",
174
- " hints.append(tag_val)\n",
175
- " if len(hints) >= 5:\n",
176
- " break\n",
177
- " \n",
178
- " hint_text = \", \".join(hints[:5]) if hints else \"none\"\n",
179
- " \n",
180
- " return f\"\"\"Given '{name}' ({hint_text}), provide:\n",
181
- "1. Full legal name (Western order if non-latin script)\n",
182
- "2. Any stage names/aliases (comma separated)\n",
183
- "3. Gender (Male/Female/Other/Unknown)\n",
184
- "4. Top 3 most likely professions from ONLY these categories:\n",
185
- " - actor\n",
186
- " - adult performer\n",
187
- " - singer/musician\n",
188
- " - model\n",
189
- " - online personality (includes streamers, cosplayers, influencers)\n",
190
- " - public figure (includes politicians, activists, journalists, authors)\n",
191
- " - voice actor/ASMR\n",
192
- " - sports professional\n",
193
- " - tv personality (includes hosts, presenters, reality TV)\n",
194
- "\n",
195
- "5. Primary country associated\n",
196
- "\n",
197
- "IMPORTANT:\n",
198
- "- Choose professions ONLY from the 9 categories above\n",
199
- "- Provide up to 3 professions, comma-separated, ordered by relevance\n",
200
- "- Be SPECIFIC: choose the most accurate category for each role\n",
201
- "- \"online personality\" includes: streamers, cosplayers, YouTubers, influencers, content creators\n",
202
- "- Use 'Unknown' when uncertain or for fictional characters/places\n",
203
- "- For multi-role people, list all relevant categories (e.g., \"actor, singer/musician, online personality\")\n",
204
- "- For country respond with one word only, for example China or Columbia\n",
205
- "- actress = actor\n",
206
- "\n",
207
- "Respond with exactly 5 numbered lines.\"\"\"\n",
208
- "\n",
209
- "# Create prompts\n",
210
- "print(\"\\nCreating prompts...\")\n",
211
- "df['prompt'] = df.apply(create_prompt, axis=1)\n",
212
- "print(\"✅ Prompts created\")\n",
213
- "\n",
214
- "# === QUERY Gemma LOCAL ===\n",
215
- "def query_gemma_local(prompt: str) -> str:\n",
216
- " \"\"\"Query Gemma locally via transformers.\"\"\"\n",
217
- " try:\n",
218
- " # Format as chat message for GEMMA\n",
219
- " messages = [\n",
220
- " {\"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",
221
- " {\"role\": \"user\", \"content\": prompt}\n",
222
- " ]\n",
223
- " \n",
224
- " # Tokenize\n",
225
- " if hasattr(tokenizer, 'apply_chat_template'):\n",
226
- " text = tokenizer.apply_chat_template(\n",
227
- " messages,\n",
228
- " tokenize=False,\n",
229
- " add_generation_prompt=True\n",
230
- " )\n",
231
- " else:\n",
232
- " # Fallback for older tokenizers\n",
233
- " text = f\"[INST] {prompt} [/INST]\"\n",
234
- " \n",
235
- " inputs = tokenizer([text], return_tensors=\"pt\", padding=True).to(device)\n",
236
- " \n",
237
- " # Generate\n",
238
- " with torch.no_grad():\n",
239
- " outputs = model.generate(\n",
240
- " **inputs,\n",
241
- " max_new_tokens=512,\n",
242
- " temperature=0.1,\n",
243
- " do_sample=True,\n",
244
- " top_p=0.9,\n",
245
- " pad_token_id=tokenizer.eos_token_id\n",
246
- " )\n",
247
- " \n",
248
- " # Decode\n",
249
- " generated_ids = outputs[0][inputs['input_ids'].shape[1]:]\n",
250
- " response = tokenizer.decode(generated_ids, skip_special_tokens=True)\n",
251
- " \n",
252
- " return response.strip()\n",
253
- " \n",
254
- " except Exception as e:\n",
255
- " print(f\"Generation error: {e}\")\n",
256
- " return None\n",
257
- "\n",
258
- "output_file = current_dir.parent / f\"data/CSV/gemma_local_annotated_POI{'_test' if TEST_MODE else ''}.csv\"\n",
259
- "index_file = current_dir.parent / \"misc/query_indicies/gemma_local_query_index.txt\"\n",
260
- "\n",
261
- "\n",
262
- "# === PARSE RESPONSE ===\n",
263
- "def parse_response(response):\n",
264
- " \"\"\"Parse Gemma response into structured fields.\"\"\"\n",
265
- " if not response:\n",
266
- " return {\n",
267
- " 'full_name': 'Unknown',\n",
268
- " 'aliases': 'Unknown',\n",
269
- " 'gender': 'Unknown',\n",
270
- " 'profession_llm': 'Unknown',\n",
271
- " 'country': 'Unknown'\n",
272
- " }\n",
273
- " \n",
274
- " # Split into lines and clean\n",
275
- " lines = [line.strip() for line in response.split('\\n') if line.strip()]\n",
276
- " \n",
277
- " # Initialize with Unknown values\n",
278
- " fields = {\n",
279
- " 'full_name': 'Unknown',\n",
280
- " 'aliases': 'Unknown',\n",
281
- " 'gender': 'Unknown',\n",
282
- " 'profession_llm': 'Unknown',\n",
283
- " 'country': 'Unknown'\n",
284
- " }\n",
285
- " \n",
286
- " # Extract information from each numbered line\n",
287
- " for line in lines:\n",
288
- " if line.startswith('1.'):\n",
289
- " fields['full_name'] = line[2:].strip()\n",
290
- " elif line.startswith('2.'):\n",
291
- " fields['aliases'] = line[2:].strip()\n",
292
- " elif line.startswith('3.'):\n",
293
- " fields['gender'] = line[2:].strip()\n",
294
- " elif line.startswith('4.'):\n",
295
- " fields['profession_llm'] = line[2:].strip()\n",
296
- " elif line.startswith('5.'):\n",
297
- " fields['country'] = line[2:].strip()\n",
298
- " \n",
299
- " return fields\n",
300
- "\n",
301
- "\n",
302
- "# === PROCESS DATA ===\n",
303
- "index_file.parent.mkdir(parents=True, exist_ok=True)\n",
304
- "\n",
305
- "# Load index\n",
306
- "current_index = 0\n",
307
- "if index_file.exists():\n",
308
- " try:\n",
309
- " current_index = int(index_file.read_text().strip())\n",
310
- " except:\n",
311
- " current_index = 0\n",
312
- "\n",
313
- "print(f\"Resuming from index {current_index}\")\n",
314
- "\n",
315
- "start_time = time.time()\n",
316
- "\n",
317
- "for i in tqdm(range(current_index, len(df)), desc=\"Gemma Local\"):\n",
318
- "\n",
319
- " prompt = df.at[i, \"prompt\"]\n",
320
- "\n",
321
- " # -------- MODEL QUERY WITH RETRIES --------\n",
322
- " response = None\n",
323
- " for attempt in range(3):\n",
324
- " response = query_gemma_local(prompt)\n",
325
- " \n",
326
- " # Valid response?\n",
327
- " if response and len(response.strip()) > 10:\n",
328
- " break\n",
329
- " \n",
330
- " print(f\"⚠️ Row {i}: Empty or invalid response, retry {attempt+1}/3\")\n",
331
- " time.sleep(0.5)\n",
332
- "\n",
333
- " # If still invalid → DO NOT overwrite previous data\n",
334
- " if not response or len(response.strip()) <= 10:\n",
335
- " print(f\"❌ Row {i}: failed after retries, not writing, not advancing index\")\n",
336
- " continue\n",
337
- "\n",
338
- " parsed = parse_response(response)\n",
339
- "\n",
340
- " # Additional safety: skip rows that parsed as all 'Unknown'\n",
341
- " if all(v == \"Unknown\" for v in parsed.values()):\n",
342
- " print(f\"❌ Row {i}: parsed as all Unknown (likely model crash); skipping.\")\n",
343
- " continue\n",
344
- "\n",
345
- " # -------- WRITE PARSED FIELDS SAFELY --------\n",
346
- " for key, value in parsed.items():\n",
347
- " df.at[i, key] = value\n",
348
- "\n",
349
- " # Advance progress ONLY after successful write\n",
350
- " current_index = i + 1\n",
351
- "\n",
352
- " # -------- GPU MEMORY CLEANUP --------\n",
353
- " if torch.cuda.is_available():\n",
354
- " torch.cuda.empty_cache()\n",
355
- " torch.cuda.synchronize()\n",
356
- "\n",
357
- " # -------- SAVE LIKE YOUR DEEPSEEK VERSION --------\n",
358
- " if (i + 1) % SAVE_INTERVAL == 0 or (i + 1) == len(df):\n",
359
- " df.to_csv(output_file, index=False)\n",
360
- " with open(index_file, \"w\") as f:\n",
361
- " f.write(str(current_index))\n",
362
- " print(f\"💾 Progress saved after row {i+1}\")\n",
363
- "\n",
364
- "# Final save\n",
365
- "df.to_csv(output_file, index=False)\n",
366
- "index_file.write_text(str(current_index))\n",
367
- "print(\"✅ Finished full dataset.\")\n"
368
- ]
369
- },
370
- {
371
- "cell_type": "code",
372
- "execution_count": null,
373
- "id": "7a1be7c0-ce54-4445-8534-bf3ab5e70197",
374
- "metadata": {},
375
- "outputs": [],
376
- "source": []
377
- }
378
- ],
379
- "metadata": {
380
- "kernelspec": {
381
- "display_name": "pm-paper",
382
- "language": "python",
383
- "name": "pm-paper"
384
- },
385
- "language_info": {
386
- "codemirror_mode": {
387
- "name": "ipython",
388
- "version": 3
389
- },
390
- "file_extension": ".py",
391
- "mimetype": "text/x-python",
392
- "name": "python",
393
- "nbconvert_exporter": "python",
394
- "pygments_lexer": "ipython3",
395
- "version": "3.11.13"
396
- }
397
- },
398
- "nbformat": 4,
399
- "nbformat_minor": 5
400
- }
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
jupyter_notebooks/MISTRAL.ipynb DELETED
The diff for this file is too large to render. See raw diff
 
jupyter_notebooks/QWEN.ipynb DELETED
@@ -1,1341 +0,0 @@
1
- {
2
- "cells": [
3
- {
4
- "cell_type": "markdown",
5
- "id": "0543d9c4-055b-49a3-a8d0-9dfb622b2b8c",
6
- "metadata": {},
7
- "source": [
8
- "# QWEN 2.5-32B Local Inference\n",
9
- "\n",
10
- "## Hardware Requirements\n",
11
- "- **GPU**: NVIDIA A100 (40GB or 80GB recommended)\n",
12
- "- **VRAM Usage**: \n",
13
- " - 8-bit quantization: ~32GB\n",
14
- " - 4-bit quantization: ~16-20GB\n",
15
- " - bfloat16 (no quantization): ~64GB\n",
16
- "- **System RAM**: 32GB minimum, 64GB recommended\n",
17
- "- **Storage**: ~65GB for model download\n",
18
- "\n",
19
- "## Configuration\n",
20
- "This notebook uses **8-bit quantization** via `bitsandbytes` for optimal performance on A100 GPUs:\n",
21
- "- Reduces VRAM usage from 64GB to ~32GB\n",
22
- "- Minimal quality degradation\n",
23
- "- Faster inference than bfloat16\n",
24
- "\n",
25
- "## Model Details\n",
26
- "- **Model**: Qwen/Qwen2.5-32B-Instruct\n",
27
- "- **Task**: Entity annotation and profession classification\n",
28
- "- **Quantization**: LLM.int8() (8-bit)\n",
29
- "- **Device**: CUDA (auto device mapping)\n",
30
- "\n",
31
- "## Dependencies\n",
32
- "Make sure to install:\n",
33
- "```bash\n",
34
- "pip install transformers>=4.35.0 bitsandbytes>=0.41.0 accelerate torch pandas tqdm\n",
35
- "```"
36
- ]
37
- },
38
- {
39
- "cell_type": "code",
40
- "execution_count": null,
41
- "id": "fe6ba282-896b-4272-b82b-ef24810732fb",
42
- "metadata": {
43
- "execution": {
44
- "iopub.execute_input": "2025-12-07T21:04:54.429449Z",
45
- "iopub.status.busy": "2025-12-07T21:04:54.429316Z"
46
- }
47
- },
48
- "outputs": [
49
- {
50
- "name": "stderr",
51
- "output_type": "stream",
52
- "text": [
53
- "/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",
54
- " from .autonotebook import tqdm as notebook_tqdm\n"
55
- ]
56
- },
57
- {
58
- "name": "stdout",
59
- "output_type": "stream",
60
- "text": [
61
- "Loading model: Qwen/Qwen2.5-32B-Instruct\n",
62
- "Cache directory: /shares/weddigen.ki.uzh/laura_wagner/phase_01/pm-paper/data/models\n",
63
- "This may take a while on first run (~65GB download)...\n",
64
- "\n",
65
- "Device: cuda\n",
66
- "Loading tokenizer...\n",
67
- "✅ Tokenizer loaded\n",
68
- "Configuring 8-bit quantization...\n",
69
- "Loading model with 8-bit quantization (this may take several minutes)...\n"
70
- ]
71
- },
72
- {
73
- "name": "stderr",
74
- "output_type": "stream",
75
- "text": [
76
- "Loading checkpoint shards: 100%|██████████| 17/17 [03:55<00:00, 13.84s/it]\n"
77
- ]
78
- },
79
- {
80
- "name": "stdout",
81
- "output_type": "stream",
82
- "text": [
83
- "✅ Model loaded with 8-bit quantization\n",
84
- "VRAM used: 32.72 GB\n",
85
- "\n",
86
- "Loading annotated CSV...\n"
87
- ]
88
- },
89
- {
90
- "name": "stderr",
91
- "output_type": "stream",
92
- "text": [
93
- "/tmp/ipykernel_1943010/494578941.py:113: DtypeWarning: Columns (53,54,55,56,57) have mixed types. Specify dtype option on import or set low_memory=False.\n",
94
- " df = pd.read_csv(output_file)\n"
95
- ]
96
- },
97
- {
98
- "name": "stdout",
99
- "output_type": "stream",
100
- "text": [
101
- "✅ Loaded professions.csv\n",
102
- "✅ Loaded profession mapping with 9 categories\n",
103
- "Loaded 50861 rows\n",
104
- "\n",
105
- "Profession categories (9):\n",
106
- " - actor\n",
107
- " - adult performer\n",
108
- " - singer/musician\n",
109
- " - model\n",
110
- " - online personality\n",
111
- " - public figure\n",
112
- " - voice actor/ASMR\n",
113
- " - sports professional\n",
114
- " - tv personality\n",
115
- "\n",
116
- "Creating prompts...\n",
117
- "✅ Prompts created\n",
118
- "Resuming from index 48500\n"
119
- ]
120
- },
121
- {
122
- "name": "stderr",
123
- "output_type": "stream",
124
- "text": [
125
- "Qwen Local: 0%| | 0/2361 [00:00<?, ?it/s]The following generation flags are not valid and may be ignored: ['temperature', 'top_p', 'top_k']. Set `TRANSFORMERS_VERBOSITY=info` for more details.\n",
126
- "Qwen Local: 0%| | 10/2361 [01:07<4:34:08, 7.00s/it]"
127
- ]
128
- },
129
- {
130
- "name": "stdout",
131
- "output_type": "stream",
132
- "text": [
133
- "💾 Progress saved after row 48510\n"
134
- ]
135
- },
136
- {
137
- "name": "stderr",
138
- "output_type": "stream",
139
- "text": [
140
- "Qwen Local: 1%| | 20/2361 [02:16<4:45:58, 7.33s/it]"
141
- ]
142
- },
143
- {
144
- "name": "stdout",
145
- "output_type": "stream",
146
- "text": [
147
- "💾 Progress saved after row 48520\n"
148
- ]
149
- },
150
- {
151
- "name": "stderr",
152
- "output_type": "stream",
153
- "text": [
154
- "Qwen Local: 1%|▏ | 30/2361 [03:24<4:52:57, 7.54s/it]"
155
- ]
156
- },
157
- {
158
- "name": "stdout",
159
- "output_type": "stream",
160
- "text": [
161
- "💾 Progress saved after row 48530\n"
162
- ]
163
- },
164
- {
165
- "name": "stderr",
166
- "output_type": "stream",
167
- "text": [
168
- "Qwen Local: 2%|▏ | 40/2361 [04:37<5:05:46, 7.90s/it]"
169
- ]
170
- },
171
- {
172
- "name": "stdout",
173
- "output_type": "stream",
174
- "text": [
175
- "💾 Progress saved after row 48540\n"
176
- ]
177
- },
178
- {
179
- "name": "stderr",
180
- "output_type": "stream",
181
- "text": [
182
- "Qwen Local: 2%|▏ | 50/2361 [05:48<4:44:32, 7.39s/it]"
183
- ]
184
- },
185
- {
186
- "name": "stdout",
187
- "output_type": "stream",
188
- "text": [
189
- "💾 Progress saved after row 48550\n"
190
- ]
191
- },
192
- {
193
- "name": "stderr",
194
- "output_type": "stream",
195
- "text": [
196
- "Qwen Local: 3%|▎ | 60/2361 [06:58<4:41:41, 7.35s/it]"
197
- ]
198
- },
199
- {
200
- "name": "stdout",
201
- "output_type": "stream",
202
- "text": [
203
- "💾 Progress saved after row 48560\n"
204
- ]
205
- },
206
- {
207
- "name": "stderr",
208
- "output_type": "stream",
209
- "text": [
210
- "Qwen Local: 3%|▎ | 70/2361 [08:12<4:52:28, 7.66s/it]"
211
- ]
212
- },
213
- {
214
- "name": "stdout",
215
- "output_type": "stream",
216
- "text": [
217
- "💾 Progress saved after row 48570\n"
218
- ]
219
- },
220
- {
221
- "name": "stderr",
222
- "output_type": "stream",
223
- "text": [
224
- "Qwen Local: 3%|▎ | 80/2361 [09:19<4:36:29, 7.27s/it]"
225
- ]
226
- },
227
- {
228
- "name": "stdout",
229
- "output_type": "stream",
230
- "text": [
231
- "💾 Progress saved after row 48580\n"
232
- ]
233
- },
234
- {
235
- "name": "stderr",
236
- "output_type": "stream",
237
- "text": [
238
- "Qwen Local: 4%|▍ | 90/2361 [10:30<4:47:42, 7.60s/it]"
239
- ]
240
- },
241
- {
242
- "name": "stdout",
243
- "output_type": "stream",
244
- "text": [
245
- "💾 Progress saved after row 48590\n"
246
- ]
247
- },
248
- {
249
- "name": "stderr",
250
- "output_type": "stream",
251
- "text": [
252
- "Qwen Local: 4%|▍ | 100/2361 [11:32<4:16:12, 6.80s/it]"
253
- ]
254
- },
255
- {
256
- "name": "stdout",
257
- "output_type": "stream",
258
- "text": [
259
- "💾 Progress saved after row 48600\n"
260
- ]
261
- },
262
- {
263
- "name": "stderr",
264
- "output_type": "stream",
265
- "text": [
266
- "Qwen Local: 5%|▍ | 110/2361 [12:41<4:31:59, 7.25s/it]"
267
- ]
268
- },
269
- {
270
- "name": "stdout",
271
- "output_type": "stream",
272
- "text": [
273
- "💾 Progress saved after row 48610\n"
274
- ]
275
- },
276
- {
277
- "name": "stderr",
278
- "output_type": "stream",
279
- "text": [
280
- "Qwen Local: 5%|▌ | 120/2361 [13:48<4:39:07, 7.47s/it]"
281
- ]
282
- },
283
- {
284
- "name": "stdout",
285
- "output_type": "stream",
286
- "text": [
287
- "💾 Progress saved after row 48620\n"
288
- ]
289
- },
290
- {
291
- "name": "stderr",
292
- "output_type": "stream",
293
- "text": [
294
- "Qwen Local: 6%|▌ | 130/2361 [14:54<4:20:42, 7.01s/it]"
295
- ]
296
- },
297
- {
298
- "name": "stdout",
299
- "output_type": "stream",
300
- "text": [
301
- "💾 Progress saved after row 48630\n"
302
- ]
303
- },
304
- {
305
- "name": "stderr",
306
- "output_type": "stream",
307
- "text": [
308
- "Qwen Local: 6%|▌ | 140/2361 [15:57<4:10:07, 6.76s/it]"
309
- ]
310
- },
311
- {
312
- "name": "stdout",
313
- "output_type": "stream",
314
- "text": [
315
- "💾 Progress saved after row 48640\n"
316
- ]
317
- },
318
- {
319
- "name": "stderr",
320
- "output_type": "stream",
321
- "text": [
322
- "Qwen Local: 6%|▋ | 150/2361 [17:06<4:33:11, 7.41s/it]"
323
- ]
324
- },
325
- {
326
- "name": "stdout",
327
- "output_type": "stream",
328
- "text": [
329
- "💾 Progress saved after row 48650\n"
330
- ]
331
- },
332
- {
333
- "name": "stderr",
334
- "output_type": "stream",
335
- "text": [
336
- "Qwen Local: 7%|▋ | 160/2361 [18:17<4:50:10, 7.91s/it]"
337
- ]
338
- },
339
- {
340
- "name": "stdout",
341
- "output_type": "stream",
342
- "text": [
343
- "💾 Progress saved after row 48660\n"
344
- ]
345
- },
346
- {
347
- "name": "stderr",
348
- "output_type": "stream",
349
- "text": [
350
- "Qwen Local: 7%|▋ | 170/2361 [19:24<4:12:48, 6.92s/it]"
351
- ]
352
- },
353
- {
354
- "name": "stdout",
355
- "output_type": "stream",
356
- "text": [
357
- "💾 Progress saved after row 48670\n"
358
- ]
359
- },
360
- {
361
- "name": "stderr",
362
- "output_type": "stream",
363
- "text": [
364
- "Qwen Local: 8%|▊ | 180/2361 [20:32<4:45:13, 7.85s/it]"
365
- ]
366
- },
367
- {
368
- "name": "stdout",
369
- "output_type": "stream",
370
- "text": [
371
- "💾 Progress saved after row 48680\n"
372
- ]
373
- },
374
- {
375
- "name": "stderr",
376
- "output_type": "stream",
377
- "text": [
378
- "Qwen Local: 8%|▊ | 190/2361 [21:42<4:17:34, 7.12s/it]"
379
- ]
380
- },
381
- {
382
- "name": "stdout",
383
- "output_type": "stream",
384
- "text": [
385
- "💾 Progress saved after row 48690\n"
386
- ]
387
- },
388
- {
389
- "name": "stderr",
390
- "output_type": "stream",
391
- "text": [
392
- "Qwen Local: 8%|▊ | 200/2361 [22:51<4:22:23, 7.29s/it]"
393
- ]
394
- },
395
- {
396
- "name": "stdout",
397
- "output_type": "stream",
398
- "text": [
399
- "💾 Progress saved after row 48700\n"
400
- ]
401
- },
402
- {
403
- "name": "stderr",
404
- "output_type": "stream",
405
- "text": [
406
- "Qwen Local: 9%|▉ | 210/2361 [23:58<4:16:32, 7.16s/it]"
407
- ]
408
- },
409
- {
410
- "name": "stdout",
411
- "output_type": "stream",
412
- "text": [
413
- "💾 Progress saved after row 48710\n"
414
- ]
415
- },
416
- {
417
- "name": "stderr",
418
- "output_type": "stream",
419
- "text": [
420
- "Qwen Local: 9%|▉ | 220/2361 [25:11<4:30:09, 7.57s/it]"
421
- ]
422
- },
423
- {
424
- "name": "stdout",
425
- "output_type": "stream",
426
- "text": [
427
- "💾 Progress saved after row 48720\n"
428
- ]
429
- },
430
- {
431
- "name": "stderr",
432
- "output_type": "stream",
433
- "text": [
434
- "Qwen Local: 10%|▉ | 230/2361 [26:13<3:56:48, 6.67s/it]"
435
- ]
436
- },
437
- {
438
- "name": "stdout",
439
- "output_type": "stream",
440
- "text": [
441
- "💾 Progress saved after row 48730\n"
442
- ]
443
- },
444
- {
445
- "name": "stderr",
446
- "output_type": "stream",
447
- "text": [
448
- "Qwen Local: 10%|█ | 240/2361 [27:21<4:19:15, 7.33s/it]"
449
- ]
450
- },
451
- {
452
- "name": "stdout",
453
- "output_type": "stream",
454
- "text": [
455
- "💾 Progress saved after row 48740\n"
456
- ]
457
- },
458
- {
459
- "name": "stderr",
460
- "output_type": "stream",
461
- "text": [
462
- "Qwen Local: 11%|█ | 250/2361 [28:30<4:25:08, 7.54s/it]"
463
- ]
464
- },
465
- {
466
- "name": "stdout",
467
- "output_type": "stream",
468
- "text": [
469
- "💾 Progress saved after row 48750\n"
470
- ]
471
- },
472
- {
473
- "name": "stderr",
474
- "output_type": "stream",
475
- "text": [
476
- "Qwen Local: 11%|█ | 260/2361 [29:33<3:52:41, 6.65s/it]"
477
- ]
478
- },
479
- {
480
- "name": "stdout",
481
- "output_type": "stream",
482
- "text": [
483
- "💾 Progress saved after row 48760\n"
484
- ]
485
- },
486
- {
487
- "name": "stderr",
488
- "output_type": "stream",
489
- "text": [
490
- "Qwen Local: 11%|█▏ | 270/2361 [30:42<4:35:01, 7.89s/it]"
491
- ]
492
- },
493
- {
494
- "name": "stdout",
495
- "output_type": "stream",
496
- "text": [
497
- "💾 Progress saved after row 48770\n"
498
- ]
499
- },
500
- {
501
- "name": "stderr",
502
- "output_type": "stream",
503
- "text": [
504
- "Qwen Local: 12%|█▏ | 280/2361 [31:51<4:24:56, 7.64s/it]"
505
- ]
506
- },
507
- {
508
- "name": "stdout",
509
- "output_type": "stream",
510
- "text": [
511
- "💾 Progress saved after row 48780\n"
512
- ]
513
- },
514
- {
515
- "name": "stderr",
516
- "output_type": "stream",
517
- "text": [
518
- "Qwen Local: 12%|█▏ | 290/2361 [33:06<4:36:14, 8.00s/it]"
519
- ]
520
- },
521
- {
522
- "name": "stdout",
523
- "output_type": "stream",
524
- "text": [
525
- "💾 Progress saved after row 48790\n"
526
- ]
527
- },
528
- {
529
- "name": "stderr",
530
- "output_type": "stream",
531
- "text": [
532
- "Qwen Local: 13%|█▎ | 300/2361 [34:16<4:15:00, 7.42s/it]"
533
- ]
534
- },
535
- {
536
- "name": "stdout",
537
- "output_type": "stream",
538
- "text": [
539
- "💾 Progress saved after row 48800\n"
540
- ]
541
- },
542
- {
543
- "name": "stderr",
544
- "output_type": "stream",
545
- "text": [
546
- "Qwen Local: 13%|█▎ | 310/2361 [35:25<4:02:07, 7.08s/it]"
547
- ]
548
- },
549
- {
550
- "name": "stdout",
551
- "output_type": "stream",
552
- "text": [
553
- "💾 Progress saved after row 48810\n"
554
- ]
555
- },
556
- {
557
- "name": "stderr",
558
- "output_type": "stream",
559
- "text": [
560
- "Qwen Local: 14%|█▎ | 320/2361 [36:33<4:10:07, 7.35s/it]"
561
- ]
562
- },
563
- {
564
- "name": "stdout",
565
- "output_type": "stream",
566
- "text": [
567
- "💾 Progress saved after row 48820\n"
568
- ]
569
- },
570
- {
571
- "name": "stderr",
572
- "output_type": "stream",
573
- "text": [
574
- "Qwen Local: 14%|█▍ | 330/2361 [37:40<4:02:47, 7.17s/it]"
575
- ]
576
- },
577
- {
578
- "name": "stdout",
579
- "output_type": "stream",
580
- "text": [
581
- "💾 Progress saved after row 48830\n"
582
- ]
583
- },
584
- {
585
- "name": "stderr",
586
- "output_type": "stream",
587
- "text": [
588
- "Qwen Local: 14%|█▍ | 340/2361 [38:55<4:12:49, 7.51s/it]"
589
- ]
590
- },
591
- {
592
- "name": "stdout",
593
- "output_type": "stream",
594
- "text": [
595
- "💾 Progress saved after row 48840\n"
596
- ]
597
- },
598
- {
599
- "name": "stderr",
600
- "output_type": "stream",
601
- "text": [
602
- "Qwen Local: 15%|█▍ | 350/2361 [40:06<4:13:06, 7.55s/it]"
603
- ]
604
- },
605
- {
606
- "name": "stdout",
607
- "output_type": "stream",
608
- "text": [
609
- "💾 Progress saved after row 48850\n"
610
- ]
611
- },
612
- {
613
- "name": "stderr",
614
- "output_type": "stream",
615
- "text": [
616
- "Qwen Local: 15%|█▌ | 360/2361 [41:18<4:18:31, 7.75s/it]"
617
- ]
618
- },
619
- {
620
- "name": "stdout",
621
- "output_type": "stream",
622
- "text": [
623
- "💾 Progress saved after row 48860\n"
624
- ]
625
- },
626
- {
627
- "name": "stderr",
628
- "output_type": "stream",
629
- "text": [
630
- "Qwen Local: 16%|█▌ | 370/2361 [42:27<4:25:14, 7.99s/it]"
631
- ]
632
- },
633
- {
634
- "name": "stdout",
635
- "output_type": "stream",
636
- "text": [
637
- "💾 Progress saved after row 48870\n"
638
- ]
639
- },
640
- {
641
- "name": "stderr",
642
- "output_type": "stream",
643
- "text": [
644
- "Qwen Local: 16%|█▌ | 380/2361 [43:37<3:58:31, 7.22s/it]"
645
- ]
646
- },
647
- {
648
- "name": "stdout",
649
- "output_type": "stream",
650
- "text": [
651
- "💾 Progress saved after row 48880\n"
652
- ]
653
- },
654
- {
655
- "name": "stderr",
656
- "output_type": "stream",
657
- "text": [
658
- "Qwen Local: 17%|█▋ | 390/2361 [44:44<3:59:45, 7.30s/it]"
659
- ]
660
- },
661
- {
662
- "name": "stdout",
663
- "output_type": "stream",
664
- "text": [
665
- "💾 Progress saved after row 48890\n"
666
- ]
667
- },
668
- {
669
- "name": "stderr",
670
- "output_type": "stream",
671
- "text": [
672
- "Qwen Local: 17%|█▋ | 400/2361 [45:53<4:12:26, 7.72s/it]"
673
- ]
674
- },
675
- {
676
- "name": "stdout",
677
- "output_type": "stream",
678
- "text": [
679
- "💾 Progress saved after row 48900\n"
680
- ]
681
- },
682
- {
683
- "name": "stderr",
684
- "output_type": "stream",
685
- "text": [
686
- "Qwen Local: 17%|█▋ | 410/2361 [47:02<4:07:58, 7.63s/it]"
687
- ]
688
- },
689
- {
690
- "name": "stdout",
691
- "output_type": "stream",
692
- "text": [
693
- "💾 Progress saved after row 48910\n"
694
- ]
695
- },
696
- {
697
- "name": "stderr",
698
- "output_type": "stream",
699
- "text": [
700
- "Qwen Local: 18%|█▊ | 420/2361 [48:14<4:12:44, 7.81s/it]"
701
- ]
702
- },
703
- {
704
- "name": "stdout",
705
- "output_type": "stream",
706
- "text": [
707
- "💾 Progress saved after row 48920\n"
708
- ]
709
- },
710
- {
711
- "name": "stderr",
712
- "output_type": "stream",
713
- "text": [
714
- "Qwen Local: 18%|█▊ | 430/2361 [49:24<4:00:51, 7.48s/it]"
715
- ]
716
- },
717
- {
718
- "name": "stdout",
719
- "output_type": "stream",
720
- "text": [
721
- "💾 Progress saved after row 48930\n"
722
- ]
723
- },
724
- {
725
- "name": "stderr",
726
- "output_type": "stream",
727
- "text": [
728
- "Qwen Local: 19%|█▊ | 440/2361 [50:37<4:26:00, 8.31s/it]"
729
- ]
730
- },
731
- {
732
- "name": "stdout",
733
- "output_type": "stream",
734
- "text": [
735
- "💾 Progress saved after row 48940\n"
736
- ]
737
- },
738
- {
739
- "name": "stderr",
740
- "output_type": "stream",
741
- "text": [
742
- "Qwen Local: 19%|█▉ | 450/2361 [51:42<3:43:29, 7.02s/it]"
743
- ]
744
- },
745
- {
746
- "name": "stdout",
747
- "output_type": "stream",
748
- "text": [
749
- "💾 Progress saved after row 48950\n"
750
- ]
751
- },
752
- {
753
- "name": "stderr",
754
- "output_type": "stream",
755
- "text": [
756
- "Qwen Local: 19%|█▉ | 460/2361 [52:49<3:33:21, 6.73s/it]"
757
- ]
758
- },
759
- {
760
- "name": "stdout",
761
- "output_type": "stream",
762
- "text": [
763
- "💾 Progress saved after row 48960\n"
764
- ]
765
- },
766
- {
767
- "name": "stderr",
768
- "output_type": "stream",
769
- "text": [
770
- "Qwen Local: 20%|█▉ | 470/2361 [54:00<3:42:24, 7.06s/it]"
771
- ]
772
- },
773
- {
774
- "name": "stdout",
775
- "output_type": "stream",
776
- "text": [
777
- "💾 Progress saved after row 48970\n"
778
- ]
779
- },
780
- {
781
- "name": "stderr",
782
- "output_type": "stream",
783
- "text": [
784
- "Qwen Local: 20%|██ | 480/2361 [55:08<3:55:25, 7.51s/it]"
785
- ]
786
- },
787
- {
788
- "name": "stdout",
789
- "output_type": "stream",
790
- "text": [
791
- "💾 Progress saved after row 48980\n"
792
- ]
793
- },
794
- {
795
- "name": "stderr",
796
- "output_type": "stream",
797
- "text": [
798
- "Qwen Local: 21%|██ | 490/2361 [56:15<3:37:41, 6.98s/it]"
799
- ]
800
- },
801
- {
802
- "name": "stdout",
803
- "output_type": "stream",
804
- "text": [
805
- "💾 Progress saved after row 48990\n"
806
- ]
807
- },
808
- {
809
- "name": "stderr",
810
- "output_type": "stream",
811
- "text": [
812
- "Qwen Local: 21%|██ | 500/2361 [57:22<3:48:29, 7.37s/it]"
813
- ]
814
- },
815
- {
816
- "name": "stdout",
817
- "output_type": "stream",
818
- "text": [
819
- "💾 Progress saved after row 49000\n"
820
- ]
821
- },
822
- {
823
- "name": "stderr",
824
- "output_type": "stream",
825
- "text": [
826
- "Qwen Local: 22%|██▏ | 510/2361 [58:30<3:45:27, 7.31s/it]"
827
- ]
828
- },
829
- {
830
- "name": "stdout",
831
- "output_type": "stream",
832
- "text": [
833
- "💾 Progress saved after row 49010\n"
834
- ]
835
- },
836
- {
837
- "name": "stderr",
838
- "output_type": "stream",
839
- "text": [
840
- "Qwen Local: 22%|██▏ | 520/2361 [59:42<3:56:27, 7.71s/it]"
841
- ]
842
- },
843
- {
844
- "name": "stdout",
845
- "output_type": "stream",
846
- "text": [
847
- "💾 Progress saved after row 49020\n"
848
- ]
849
- },
850
- {
851
- "name": "stderr",
852
- "output_type": "stream",
853
- "text": [
854
- "Qwen Local: 22%|██▏ | 530/2361 [1:00:50<3:42:52, 7.30s/it]"
855
- ]
856
- },
857
- {
858
- "name": "stdout",
859
- "output_type": "stream",
860
- "text": [
861
- "💾 Progress saved after row 49030\n"
862
- ]
863
- },
864
- {
865
- "name": "stderr",
866
- "output_type": "stream",
867
- "text": [
868
- "Qwen Local: 23%|██▎ | 540/2361 [1:01:57<3:33:27, 7.03s/it]"
869
- ]
870
- },
871
- {
872
- "name": "stdout",
873
- "output_type": "stream",
874
- "text": [
875
- "💾 Progress saved after row 49040\n"
876
- ]
877
- },
878
- {
879
- "name": "stderr",
880
- "output_type": "stream",
881
- "text": [
882
- "Qwen Local: 23%|██▎ | 550/2361 [1:03:06<3:39:19, 7.27s/it]"
883
- ]
884
- },
885
- {
886
- "name": "stdout",
887
- "output_type": "stream",
888
- "text": [
889
- "💾 Progress saved after row 49050\n"
890
- ]
891
- },
892
- {
893
- "name": "stderr",
894
- "output_type": "stream",
895
- "text": [
896
- "Qwen Local: 24%|██▎ | 560/2361 [1:04:13<3:30:51, 7.02s/it]"
897
- ]
898
- },
899
- {
900
- "name": "stdout",
901
- "output_type": "stream",
902
- "text": [
903
- "💾 Progress saved after row 49060\n"
904
- ]
905
- },
906
- {
907
- "name": "stderr",
908
- "output_type": "stream",
909
- "text": [
910
- "Qwen Local: 24%|██▍ | 570/2361 [1:05:23<3:43:28, 7.49s/it]"
911
- ]
912
- },
913
- {
914
- "name": "stdout",
915
- "output_type": "stream",
916
- "text": [
917
- "💾 Progress saved after row 49070\n"
918
- ]
919
- },
920
- {
921
- "name": "stderr",
922
- "output_type": "stream",
923
- "text": [
924
- "Qwen Local: 24%|██▍ | 571/2361 [1:05:30<3:36:42, 7.26s/it]"
925
- ]
926
- }
927
- ],
928
- "source": [
929
- "import pandas as pd\n",
930
- "import json\n",
931
- "import time\n",
932
- "import re\n",
933
- "from pathlib import Path\n",
934
- "from tqdm import tqdm\n",
935
- "import torch\n",
936
- "from transformers import AutoModelForCausalLM, AutoTokenizer, BitsAndBytesConfig\n",
937
- "import signal\n",
938
- "from contextlib import contextmanager\n",
939
- "\n",
940
- "current_dir = Path.cwd()\n",
941
- "input_file = current_dir.parent / \"data/CSV/model_adapter/real_person_adapter_step_02_NER.csv\"\n",
942
- "professions_file = current_dir.parent / \"misc/lists/professions.csv\"\n",
943
- "professions_mapped_file = current_dir.parent / \"misc/lists/professions_mapped.csv\"\n",
944
- "# === PROCESS DATA ===\n",
945
- "\n",
946
- "\n",
947
- "# === CONFIGURATION ===\n",
948
- "TEST_MODE = False\n",
949
- "TEST_SIZE = 100\n",
950
- "MAX_ROWS = 50862\n",
951
- "SAVE_INTERVAL = 10\n",
952
- "\n",
953
- "\n",
954
- "index_file = current_dir.parent / \"misc/query_indicies/qwen_local_query_index.txt\"\n",
955
- "output_file = current_dir.parent / f\"data/CSV/qwen_local_annotated_POI{'_test' if TEST_MODE else ''}.csv\"\n",
956
- "\n",
957
- "# Model settings\n",
958
- "MODEL_NAME = \"Qwen/Qwen2.5-32B-Instruct\"\n",
959
- "#MODEL_NAME = \"Qwen/Qwen2.5-14B-Instruct\"\n",
960
- "#MODEL_NAME = \"Qwen/Qwen3-235B-A22B-Instruct-2507-FP8\"\n",
961
- "#MODEL_NAME = \"mistralai/Mistral-7B-Instruct-v0.3\"\n",
962
- "CACHE_DIR = current_dir.parent / \"data/models\"\n",
963
- "CACHE_DIR.mkdir(parents=True, exist_ok=True)\n",
964
- "\n",
965
- "# Define the SPECIFIC profession categories\n",
966
- "PROFESSION_CATEGORIES = [\n",
967
- " \"actor\",\n",
968
- " \"adult performer\",\n",
969
- " \"singer/musician\",\n",
970
- " \"model\",\n",
971
- " \"online personality\",\n",
972
- " \"public figure\",\n",
973
- " \"voice actor/ASMR\",\n",
974
- " \"sports professional\",\n",
975
- " \"tv personality\"\n",
976
- "]\n",
977
- "\n",
978
- "# === LOAD MODEL ===\n",
979
- "print(f\"Loading model: {MODEL_NAME}\")\n",
980
- "print(f\"Cache directory: {CACHE_DIR}\")\n",
981
- "print(f\"This may take a while on first run (~65GB download)...\\n\")\n",
982
- "\n",
983
- "# Check GPU availability\n",
984
- "device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n",
985
- "print(f\"Device: {device}\")\n",
986
- "\n",
987
- "if device == \"cpu\":\n",
988
- " print(\"⚠️ WARNING: No GPU detected! Inference will be VERY slow.\")\n",
989
- " print(\" Consider using a GPU or reducing model size.\")\n",
990
- "\n",
991
- "# Load tokenizer\n",
992
- "print(\"Loading tokenizer...\")\n",
993
- "try:\n",
994
- " tokenizer = AutoTokenizer.from_pretrained(\n",
995
- " MODEL_NAME,\n",
996
- " cache_dir=str(CACHE_DIR),\n",
997
- " use_fast=True\n",
998
- " )\n",
999
- "except Exception as e:\n",
1000
- " print(f\"Failed with use_fast=True, trying use_fast=False...\")\n",
1001
- " tokenizer = AutoTokenizer.from_pretrained(\n",
1002
- " MODEL_NAME,\n",
1003
- " cache_dir=str(CACHE_DIR),\n",
1004
- " use_fast=False\n",
1005
- " )\n",
1006
- "\n",
1007
- "# Ensure pad token is set\n",
1008
- "if tokenizer.pad_token is None:\n",
1009
- " tokenizer.pad_token = tokenizer.eos_token\n",
1010
- "\n",
1011
- "print(\"✅ Tokenizer loaded\")\n",
1012
- "\n",
1013
- "# Configure 8-bit quantization for A100\n",
1014
- "print(\"Configuring 8-bit quantization...\")\n",
1015
- "quantization_config = BitsAndBytesConfig(\n",
1016
- " load_in_8bit=True,\n",
1017
- " llm_int8_threshold=6.0,\n",
1018
- " llm_int8_has_fp16_weight=False\n",
1019
- ")\n",
1020
- "\n",
1021
- "# Load model with 8-bit quantization\n",
1022
- "print(\"Loading model with 8-bit quantization (this may take several minutes)...\")\n",
1023
- "model = AutoModelForCausalLM.from_pretrained(\n",
1024
- " MODEL_NAME,\n",
1025
- " cache_dir=str(CACHE_DIR),\n",
1026
- " quantization_config=quantization_config,\n",
1027
- " device_map=\"auto\",\n",
1028
- " trust_remote_code=False\n",
1029
- ")\n",
1030
- "model.eval()\n",
1031
- "print(\"✅ Model loaded with 8-bit quantization\")\n",
1032
- "\n",
1033
- "# Check VRAM usage\n",
1034
- "if torch.cuda.is_available():\n",
1035
- " vram_gb = torch.cuda.max_memory_allocated() / 1024**3\n",
1036
- " print(f\"VRAM used: {vram_gb:.2f} GB\\n\")\n",
1037
- "\n",
1038
- "# === LOAD DATA ===\n",
1039
- "if output_file.exists():\n",
1040
- " print(\"Loading annotated CSV...\")\n",
1041
- " df = pd.read_csv(output_file)\n",
1042
- "else:\n",
1043
- " print(\"Loading raw input CSV...\")\n",
1044
- " df = pd.read_csv(input_file)\n",
1045
- "\n",
1046
- "\n",
1047
- "# Try to load profession mapping files\n",
1048
- "try:\n",
1049
- " professions_df = pd.read_csv(professions_file)\n",
1050
- " print(f\"✅ Loaded professions.csv\")\n",
1051
- "except:\n",
1052
- " print(\"⚠️ Warning: professions.csv not found\")\n",
1053
- "\n",
1054
- "try:\n",
1055
- " prof_mapped_df = pd.read_csv(professions_mapped_file)\n",
1056
- " print(f\"✅ Loaded profession mapping with {len(prof_mapped_df)} categories\")\n",
1057
- "except:\n",
1058
- " print(\"⚠️ Warning: professions_mapped.csv not found, using default categories\")\n",
1059
- "\n",
1060
- "profession_str = \", \".join(PROFESSION_CATEGORIES)\n",
1061
- "\n",
1062
- "print(f\"Loaded {len(df)} rows\")\n",
1063
- "print(f\"\\nProfession categories ({len(PROFESSION_CATEGORIES)}):\")\n",
1064
- "for cat in PROFESSION_CATEGORIES:\n",
1065
- " print(f\" - {cat}\")\n",
1066
- "\n",
1067
- "if TEST_MODE:\n",
1068
- " print(f\"\\nRunning in TEST MODE with {TEST_SIZE} samples\")\n",
1069
- " df = df.head(TEST_SIZE).copy()\n",
1070
- "elif MAX_ROWS:\n",
1071
- " df = df.head(MAX_ROWS).copy()\n",
1072
- "\n",
1073
- "# === CREATE PROMPTS (OPTIMIZED FOR CLEAN OUTPUTS) ===\n",
1074
- "def create_prompt(row):\n",
1075
- " \"\"\"Create prompt for Qwen annotation with strict formatting requirements.\"\"\"\n",
1076
- " name = row['real_name'] if pd.notna(row.get('real_name')) else row.get('name', '')\n",
1077
- " \n",
1078
- " # Gather hints\n",
1079
- " hints = []\n",
1080
- " if pd.notna(row.get('likely_profession')):\n",
1081
- " hints.append(str(row['likely_profession']))\n",
1082
- " if pd.notna(row.get('likely_nationality')):\n",
1083
- " hints.append(str(row['likely_nationality']))\n",
1084
- " if pd.notna(row.get('likely_country')):\n",
1085
- " hints.append(str(row['likely_country']))\n",
1086
- " \n",
1087
- " # Add tags if we don't have enough hints\n",
1088
- " if len(hints) < 3:\n",
1089
- " for i in range(1, 8):\n",
1090
- " tag_col = f'tag_{i}'\n",
1091
- " if tag_col in row and pd.notna(row[tag_col]):\n",
1092
- " tag_val = str(row[tag_col])\n",
1093
- " if tag_val not in hints:\n",
1094
- " hints.append(tag_val)\n",
1095
- " if len(hints) >= 5:\n",
1096
- " break\n",
1097
- " \n",
1098
- " hint_text = \", \".join(hints[:5]) if hints else \"none\"\n",
1099
- " \n",
1100
- " return f\"\"\"Extract information about '{name}' ({hint_text}).\n",
1101
- "\n",
1102
- "Respond with EXACTLY 5 numbered lines. Each line must contain ONLY the value, no labels or extra text.\n",
1103
- "\n",
1104
- "FORMAT REQUIREMENTS:\n",
1105
- "1. Full legal name in Western order (first last). VALUE ONLY.\n",
1106
- "2. Stage names/aliases, comma-separated. If none, write \"None\". VALUE ONLY.\n",
1107
- "3. Gender: MUST be exactly one word: Male, Female, Other, or Unknown. VALUE ONLY.\n",
1108
- "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",
1109
- "5. Primary country: Country name only (e.g., \"China\", \"United States\", \"Colombia\"). VALUE ONLY.\n",
1110
- "\n",
1111
- "RULES:\n",
1112
- "- Professions MUST match the exact categories listed (actress = actor)\n",
1113
- "- \"online personality\" includes streamers, cosplayers, YouTubers, influencers\n",
1114
- "- \"public figure\" includes politicians, activists, journalists, authors\n",
1115
- "- Use \"Unknown\" when uncertain or for fictional characters\n",
1116
- "- NO explanatory text, NO labels like \"Gender:\", NO prefixes\n",
1117
- "- For multi-role people, list up to 3 categories by relevance\n",
1118
- "\n",
1119
- "EXAMPLE FORMAT:\n",
1120
- "1. Taylor Swift\n",
1121
- "2. None\n",
1122
- "3. Female\n",
1123
- "4. singer/musician, public figure\n",
1124
- "5. United States\"\"\"\n",
1125
- "\n",
1126
- "# Create prompts\n",
1127
- "print(\"\\nCreating prompts...\")\n",
1128
- "df['prompt'] = df.apply(create_prompt, axis=1)\n",
1129
- "print(\"✅ Prompts created\")\n",
1130
- "\n",
1131
- "@contextmanager\n",
1132
- "def timeout(duration):\n",
1133
- " def handler(signum, frame):\n",
1134
- " raise TimeoutError(\"Operation timed out\")\n",
1135
- " \n",
1136
- " # Set the signal handler and alarm\n",
1137
- " signal.signal(signal.SIGALRM, handler)\n",
1138
- " signal.alarm(duration)\n",
1139
- " try:\n",
1140
- " yield\n",
1141
- " finally:\n",
1142
- " signal.alarm(0) # Disable the alarm\n",
1143
- "\n",
1144
- "\n",
1145
- "def query_qwen_local(prompt: str) -> str:\n",
1146
- " \"\"\"Query Qwen locally via transformers.\"\"\"\n",
1147
- " try:\n",
1148
- " # Format as chat message for Qwen with strict system prompt\n",
1149
- " messages = [\n",
1150
- " {\"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",
1151
- " {\"role\": \"user\", \"content\": prompt}\n",
1152
- " ]\n",
1153
- " \n",
1154
- " # Tokenize\n",
1155
- " if hasattr(tokenizer, 'apply_chat_template'):\n",
1156
- " text = tokenizer.apply_chat_template(\n",
1157
- " messages,\n",
1158
- " tokenize=False,\n",
1159
- " add_generation_prompt=True\n",
1160
- " )\n",
1161
- " else:\n",
1162
- " # Fallback for older tokenizers\n",
1163
- " text = f\"[INST] {prompt} [/INST]\"\n",
1164
- " \n",
1165
- " inputs = tokenizer([text], return_tensors=\"pt\", padding=True).to(device)\n",
1166
- " \n",
1167
- " # Generate with timeout\n",
1168
- " with timeout(60):\n",
1169
- " with torch.no_grad():\n",
1170
- " outputs = model.generate(\n",
1171
- " **inputs,\n",
1172
- " max_new_tokens=100,\n",
1173
- " temperature=0.1,\n",
1174
- " do_sample=False,\n",
1175
- " pad_token_id=tokenizer.eos_token_id\n",
1176
- " )\n",
1177
- " \n",
1178
- " # Decode\n",
1179
- " generated_ids = outputs[0][inputs['input_ids'].shape[1]:]\n",
1180
- " response = tokenizer.decode(generated_ids, skip_special_tokens=True)\n",
1181
- " \n",
1182
- " return response.strip()\n",
1183
- " \n",
1184
- " except TimeoutError:\n",
1185
- " print(f\"[ERROR] Generation timed out after 60 seconds\")\n",
1186
- " return None\n",
1187
- " except Exception as e:\n",
1188
- " print(f\"Generation error: {e}\")\n",
1189
- " import traceback\n",
1190
- " traceback.print_exc()\n",
1191
- " return None\n",
1192
- "\n",
1193
- " \n",
1194
- "# === PARSE RESPONSE WITH CLEANING ===\n",
1195
- "def parse_response(response):\n",
1196
- " \"\"\"Parse Qwen response into structured fields with cleaning.\"\"\"\n",
1197
- " if not response:\n",
1198
- " return {\n",
1199
- " 'full_name': 'Unknown',\n",
1200
- " 'aliases': 'Unknown',\n",
1201
- " 'gender': 'Unknown',\n",
1202
- " 'profession_llm': 'Unknown',\n",
1203
- " 'country': 'Unknown'\n",
1204
- " }\n",
1205
- " \n",
1206
- " # Split into lines and clean\n",
1207
- " lines = [line.strip() for line in response.split('\\n') if line.strip()]\n",
1208
- " \n",
1209
- " # Initialize with Unknown values\n",
1210
- " fields = {\n",
1211
- " 'full_name': 'Unknown',\n",
1212
- " 'aliases': 'Unknown',\n",
1213
- " 'gender': 'Unknown',\n",
1214
- " 'profession_llm': 'Unknown',\n",
1215
- " 'country': 'Unknown'\n",
1216
- " }\n",
1217
- " \n",
1218
- " # Extract information from each numbered line\n",
1219
- " for line in lines:\n",
1220
- " if line.startswith('1.'):\n",
1221
- " fields['full_name'] = line[2:].strip()\n",
1222
- " elif line.startswith('2.'):\n",
1223
- " fields['aliases'] = line[2:].strip()\n",
1224
- " elif line.startswith('3.'):\n",
1225
- " # Clean gender field - remove any labels\n",
1226
- " gender_raw = line[2:].strip()\n",
1227
- " # Remove common prefixes\n",
1228
- " gender_raw = re.sub(r'^(Gender:|gender:)\\s*', '', gender_raw, flags=re.IGNORECASE)\n",
1229
- " # Extract just the gender word\n",
1230
- " gender_match = re.search(r'\\b(Male|Female|Other|Unknown)\\b', gender_raw, re.IGNORECASE)\n",
1231
- " fields['gender'] = gender_match.group(1).capitalize() if gender_match else gender_raw\n",
1232
- " elif line.startswith('4.'):\n",
1233
- " fields['profession_llm'] = line[2:].strip()\n",
1234
- " elif line.startswith('5.'):\n",
1235
- " # Clean country field - remove any labels\n",
1236
- " country_raw = line[2:].strip()\n",
1237
- " # Remove common prefixes like \"Primary country:\", \"Country:\", etc.\n",
1238
- " country_raw = re.sub(r'^(Primary\\s+)?(associated\\s+)?country:\\s*', '', country_raw, flags=re.IGNORECASE)\n",
1239
- " fields['country'] = country_raw\n",
1240
- " \n",
1241
- " return fields\n",
1242
- "\n",
1243
- "# === PROCESS DATA ===\n",
1244
- "index_file.parent.mkdir(parents=True, exist_ok=True)\n",
1245
- "\n",
1246
- "# Load index\n",
1247
- "current_index = 0\n",
1248
- "if index_file.exists():\n",
1249
- " try:\n",
1250
- " current_index = int(index_file.read_text().strip())\n",
1251
- " except:\n",
1252
- " current_index = 0\n",
1253
- "\n",
1254
- "print(f\"Resuming from index {current_index}\")\n",
1255
- "\n",
1256
- "start_time = time.time()\n",
1257
- "\n",
1258
- "for i in tqdm(range(current_index, len(df)), desc=\"Qwen Local\"):\n",
1259
- "\n",
1260
- " prompt = df.at[i, \"prompt\"]\n",
1261
- "\n",
1262
- " # -------- MODEL QUERY WITH RETRIES --------\n",
1263
- " response = None\n",
1264
- " for attempt in range(3):\n",
1265
- " response = query_qwen_local(prompt)\n",
1266
- " \n",
1267
- " # Valid response?\n",
1268
- " if response and len(response.strip()) > 10:\n",
1269
- " break\n",
1270
- " \n",
1271
- " print(f\"⚠️ Row {i}: Empty or invalid response, retry {attempt+1}/3\")\n",
1272
- " time.sleep(0.5)\n",
1273
- "\n",
1274
- " # If still invalid → DO NOT overwrite previous data\n",
1275
- " if not response or len(response.strip()) <= 10:\n",
1276
- " print(f\"❌ Row {i}: failed after retries, not writing, not advancing index\")\n",
1277
- " continue\n",
1278
- "\n",
1279
- " parsed = parse_response(response)\n",
1280
- "\n",
1281
- " # Additional safety: skip rows that parsed as all 'Unknown'\n",
1282
- " if all(v == \"Unknown\" for v in parsed.values()):\n",
1283
- " print(f\"❌ Row {i}: parsed as all Unknown (likely model crash); skipping.\")\n",
1284
- " continue\n",
1285
- "\n",
1286
- " # -------- WRITE PARSED FIELDS SAFELY --------\n",
1287
- " for key, value in parsed.items():\n",
1288
- " df.at[i, key] = value\n",
1289
- "\n",
1290
- " # Advance progress ONLY after successful write\n",
1291
- " current_index = i + 1\n",
1292
- "\n",
1293
- " # -------- GPU MEMORY CLEANUP --------\n",
1294
- " if torch.cuda.is_available():\n",
1295
- " torch.cuda.empty_cache()\n",
1296
- " torch.cuda.synchronize()\n",
1297
- "\n",
1298
- " # -------- SAVE LIKE YOUR DEEPSEEK VERSION --------\n",
1299
- " if (i + 1) % SAVE_INTERVAL == 0 or (i + 1) == len(df):\n",
1300
- " df.to_csv(output_file, index=False)\n",
1301
- " with open(index_file, \"w\") as f:\n",
1302
- " f.write(str(current_index))\n",
1303
- " print(f\"💾 Progress saved after row {i+1}\")\n",
1304
- "\n",
1305
- "# Final save\n",
1306
- "df.to_csv(output_file, index=False)\n",
1307
- "index_file.write_text(str(current_index))\n",
1308
- "print(\"✅ Finished full dataset.\")"
1309
- ]
1310
- },
1311
- {
1312
- "cell_type": "code",
1313
- "execution_count": null,
1314
- "id": "d9c7deb9-847a-472d-8055-f93dbfa6aa2e",
1315
- "metadata": {},
1316
- "outputs": [],
1317
- "source": []
1318
- }
1319
- ],
1320
- "metadata": {
1321
- "kernelspec": {
1322
- "display_name": "pm-paper",
1323
- "language": "python",
1324
- "name": "pm-paper"
1325
- },
1326
- "language_info": {
1327
- "codemirror_mode": {
1328
- "name": "ipython",
1329
- "version": 3
1330
- },
1331
- "file_extension": ".py",
1332
- "mimetype": "text/x-python",
1333
- "name": "python",
1334
- "nbconvert_exporter": "python",
1335
- "pygments_lexer": "ipython3",
1336
- "version": "3.11.13"
1337
- }
1338
- },
1339
- "nbformat": 4,
1340
- "nbformat_minor": 5
1341
- }
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
jupyter_notebooks/Section_2-3-4_Bloomz_query.ipynb DELETED
@@ -1,350 +0,0 @@
1
- {
2
- "cells": [
3
- {
4
- "cell_type": "code",
5
- "execution_count": 13,
6
- "id": "3b87c378-241e-41ab-be6e-84222594f22f",
7
- "metadata": {
8
- "execution": {
9
- "iopub.execute_input": "2025-11-28T13:07:47.025752Z",
10
- "iopub.status.busy": "2025-11-28T13:07:47.025552Z",
11
- "iopub.status.idle": "2025-11-28T13:07:47.128799Z",
12
- "shell.execute_reply": "2025-11-28T13:07:47.128306Z",
13
- "shell.execute_reply.started": "2025-11-28T13:07:47.025736Z"
14
- }
15
- },
16
- "outputs": [],
17
- "source": [
18
- "import pandas as pd\n",
19
- "import json\n",
20
- "import time\n",
21
- "import re\n",
22
- "from pathlib import Path\n",
23
- "from tqdm import tqdm\n",
24
- "import torch\n",
25
- "from transformers import AutoModelForCausalLM, AutoTokenizer\n",
26
- "from datetime import datetime\n",
27
- "\n",
28
- "# Import is used for pd.notna() and pd.isna() checks\n",
29
- "\n",
30
- "current_dir = Path.cwd()\n",
31
- "input_file = current_dir.parent / \"data/CSV/model_adapter/real_person_adapter_step_02_NER.csv\"\n",
32
- "\n",
33
- "# === CONFIGURATION ===\n",
34
- "TEST_MODE = True\n",
35
- "TEST_SIZE = 10\n",
36
- "MAX_ROWS = 20000\n",
37
- "SAVE_INTERVAL = 20\n",
38
- "\n",
39
- "# Model settings - BLOOMZ (BigScience - European consortium)\n",
40
- "MODEL_NAME = \"bigscience/bloomz-7b1\" # Largest instruction-tuned BLOOM model\n",
41
- "CACHE_DIR = current_dir.parent / \"data/models\"\n",
42
- "CACHE_DIR.mkdir(parents=True, exist_ok=True)\n",
43
- "\n",
44
- "PROFESSION_CATEGORIES = [\n",
45
- " \"actor\", \"adult performer\", \"singer/musician\", \"model\",\n",
46
- " \"online personality\", \"public figure\", \"voice actor/ASMR\",\n",
47
- " \"sports professional\", \"tv personality\"\n",
48
- "]\n",
49
- "\n",
50
- "# === LOAD MODEL ===\n",
51
- "print(f\"Loading model: {MODEL_NAME}\")\n",
52
- "print(f\"Cache directory: {CACHE_DIR}\")\n",
53
- "print(f\"This may take a while on first run (~14GB download)...\\n\")\n",
54
- "\n",
55
- "# Check GPU availability\n",
56
- "device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n",
57
- "print(f\"Device: {device}\")\n",
58
- "\n",
59
- "if device == \"cpu\":\n",
60
- " print(\"⚠️ WARNING: No GPU detected! Inference will be VERY slow.\")\n",
61
- " print(\" Consider using a GPU or reducing model size.\")\n",
62
- "\n",
63
- "# Load tokenizer\n",
64
- "print(\"Loading tokenizer...\")\n",
65
- "try:\n",
66
- " tokenizer = AutoTokenizer.from_pretrained(\n",
67
- " MODEL_NAME,\n",
68
- " cache_dir=str(CACHE_DIR)\n",
69
- " )\n",
70
- " print(\"✅ Tokenizer loaded\")\n",
71
- "except Exception as e:\n",
72
- " print(f\"❌ Error loading tokenizer: {e}\")\n",
73
- " raise\n",
74
- "\n",
75
- "# Ensure pad token is set\n",
76
- "if tokenizer.pad_token is None:\n",
77
- " tokenizer.pad_token = tokenizer.eos_token\n",
78
- " print(f\"Set pad_token to eos_token: {tokenizer.eos_token}\")\n",
79
- "\n",
80
- "# Load model with optimizations\n",
81
- "print(\"Loading model (this may take several minutes)...\")\n",
82
- "try:\n",
83
- " model = AutoModelForCausalLM.from_pretrained(\n",
84
- " MODEL_NAME,\n",
85
- " cache_dir=str(CACHE_DIR),\n",
86
- " torch_dtype=torch.bfloat16, # Use BF16 for efficiency\n",
87
- " device_map=\"auto\", # Automatically distribute across GPUs\n",
88
- " low_cpu_mem_usage=True # Optimize memory usage\n",
89
- " )\n",
90
- " model.eval() # Set to evaluation mode\n",
91
- " print(\"✅ Model loaded\")\n",
92
- "except Exception as e:\n",
93
- " print(f\"❌ Error loading model: {e}\")\n",
94
- " raise\n",
95
- "\n",
96
- "# Check VRAM usage\n",
97
- "if torch.cuda.is_available():\n",
98
- " vram_gb = torch.cuda.max_memory_allocated() / 1024**3\n",
99
- " print(f\"VRAM used: {vram_gb:.2f} GB\\n\")\n",
100
- "\n",
101
- "# === LOAD DATA ===\n",
102
- "df = pd.read_csv(input_file)\n",
103
- "print(f\"Loaded {len(df)} rows\")\n",
104
- "\n",
105
- "if TEST_MODE:\n",
106
- " print(f\"Running in TEST MODE with {TEST_SIZE} samples\")\n",
107
- " df = df.head(TEST_SIZE).copy()\n",
108
- "elif MAX_ROWS:\n",
109
- " df = df.head(MAX_ROWS).copy()\n",
110
- "\n",
111
- "# === CREATE PROMPT (Exact DeepSeek style) ===\n",
112
- "# === CREATE PROMPT (BLOOMZ-optimized - Complete format) ===\n",
113
- "def create_prompt(row):\n",
114
- " \"\"\"Create prompt optimized for BLOOMZ with complete examples.\"\"\"\n",
115
- " name = row.get('real_name', row.get('name', ''))\n",
116
- " if pd.isna(name):\n",
117
- " name = row.get('name', '')\n",
118
- "\n",
119
- " # Gather hints\n",
120
- " hints = []\n",
121
- " if pd.notna(row.get('likely_profession')):\n",
122
- " hints.append(str(row['likely_profession']))\n",
123
- " if pd.notna(row.get('likely_nationality')):\n",
124
- " hints.append(str(row['likely_nationality']))\n",
125
- " if pd.notna(row.get('likely_country')):\n",
126
- " hints.append(str(row['likely_country']))\n",
127
- "\n",
128
- " # Add tags if we don't have enough hints\n",
129
- " if len(hints) < 3:\n",
130
- " for i in range(1, 8):\n",
131
- " tag_col = f'tag_{i}'\n",
132
- " if tag_col in row and pd.notna(row[tag_col]):\n",
133
- " tag_val = str(row[tag_col])\n",
134
- " if tag_val not in hints:\n",
135
- " hints.append(tag_val)\n",
136
- " if len(hints) >= 5:\n",
137
- " break\n",
138
- "\n",
139
- " hint_text = \", \".join(hints[:5]) if hints else \"none\"\n",
140
- "\n",
141
- " # BLOOMZ needs complete format shown - don't end with \"1.\"\n",
142
- " return f\"\"\"Task: Extract person information in 5 numbered lines.\n",
143
- "\n",
144
- "Format:\n",
145
- "1. Full legal name\n",
146
- "2. Stage names/aliases\n",
147
- "3. Gender\n",
148
- "4. Professions (from: actor, adult performer, singer/musician, model, online personality, public figure, voice actor/ASMR, sports professional, tv personality)\n",
149
- "5. Country\n",
150
- "\n",
151
- "Example:\n",
152
- "Person: Taylor Swift (singer, American, pop music)\n",
153
- "Answer:\n",
154
- "1. Taylor Alison Swift\n",
155
- "2. Taylor Swift\n",
156
- "3. Female\n",
157
- "4. singer/musician, online personality\n",
158
- "5. United States\n",
159
- "\n",
160
- "Person: {name} ({hint_text})\n",
161
- "Answer:\n",
162
- "\"\"\"\n",
163
- "\n",
164
- "# === QUERY BLOOMZ LOCAL (Fixed generation parameters) ===\n",
165
- "def query_bloomz_local(prompt: str) -> str:\n",
166
- " \"\"\"Query BLOOMZ-7B1 locally via transformers, return raw response string.\"\"\"\n",
167
- " try:\n",
168
- " # Store full prompt for logging\n",
169
- " query_bloomz_local.last_full_prompt = prompt\n",
170
- "\n",
171
- " # Tokenize WITHOUT truncation\n",
172
- " inputs = tokenizer(\n",
173
- " prompt,\n",
174
- " return_tensors=\"pt\",\n",
175
- " truncation=False\n",
176
- " ).to(device)\n",
177
- " \n",
178
- " # Store the input length for proper extraction later\n",
179
- " input_length = inputs['input_ids'].shape[1]\n",
180
- "\n",
181
- " # Generate with parameters optimized for BLOOMZ completion\n",
182
- " with torch.no_grad():\n",
183
- " outputs = model.generate(\n",
184
- " **inputs,\n",
185
- " max_new_tokens=300, # Increased to allow full completion\n",
186
- " min_new_tokens=50, # Force minimum generation length\n",
187
- " temperature=0.5, # Balanced creativity\n",
188
- " do_sample=True,\n",
189
- " top_p=0.9,\n",
190
- " top_k=50,\n",
191
- " repetition_penalty=1.15,\n",
192
- " pad_token_id=tokenizer.eos_token_id,\n",
193
- " eos_token_id=tokenizer.eos_token_id,\n",
194
- " early_stopping=False, # Changed to False - don't stop early!\n",
195
- " num_beams=1 # Greedy-like but with sampling\n",
196
- " )\n",
197
- "\n",
198
- " # Extract only the NEW generated tokens\n",
199
- " generated_ids = outputs[0][input_length:]\n",
200
- " generated_text = tokenizer.decode(generated_ids, skip_special_tokens=True)\n",
201
- "\n",
202
- " # Debug output (show first 5 now to see more patterns)\n",
203
- " if not hasattr(query_bloomz_local, 'debug_count'):\n",
204
- " query_bloomz_local.debug_count = 0\n",
205
- "\n",
206
- " if query_bloomz_local.debug_count < 5:\n",
207
- " print(f\"\\n📝 BLOOMZ Debug #{query_bloomz_local.debug_count + 1}:\")\n",
208
- " print(f\"Prompt tokens: {input_length}\")\n",
209
- " print(f\"Generated tokens: {len(generated_ids)}\")\n",
210
- " print(f\"Full generation:\\n{generated_text}\")\n",
211
- " print(f\"{'='*60}\\n\")\n",
212
- " query_bloomz_local.debug_count += 1\n",
213
- "\n",
214
- " return generated_text.strip()\n",
215
- "\n",
216
- " except Exception as e:\n",
217
- " print(f\"Error querying BLOOMZ: {e}\")\n",
218
- " import traceback\n",
219
- " traceback.print_exc()\n",
220
- " query_bloomz_local.last_full_prompt = f\"ERROR: {e}\"\n",
221
- " return None\n",
222
- "\n",
223
- " \n",
224
- "# === PROCESS ===\n",
225
- "output_file = current_dir.parent / f\"data/CSV/bloomz_annotated_POI{'_test' if TEST_MODE else ''}.csv\"\n",
226
- "index_file = current_dir.parent / \"misc/bloomz_query_index.txt\"\n",
227
- "\n",
228
- "current_index = 0\n",
229
- "if index_file.exists():\n",
230
- " with open(index_file) as f:\n",
231
- " current_index = int(f.read().strip())\n",
232
- " print(f\"Resuming from index {current_index}\")\n",
233
- "\n",
234
- "# Initialize columns (same as DeepSeek)\n",
235
- "for col in ['full_name', 'gender', 'profession_llm', 'country', 'aliases']:\n",
236
- " if col not in df.columns:\n",
237
- " df[col] = 'Unknown'\n",
238
- "\n",
239
- "# Create prompts for all rows (same as DeepSeek)\n",
240
- "print(\"Creating prompts...\")\n",
241
- "df['prompt'] = df.apply(create_prompt, axis=1)\n",
242
- "\n",
243
- "print(f\"\\nAnnotating with BLOOMZ-7B1 LOCAL - rows {current_index} to {len(df)}...\")\n",
244
- "print(f\"Model: {MODEL_NAME}\")\n",
245
- "print(f\"This may take a while...\\n\")\n",
246
- "\n",
247
- "try:\n",
248
- " start_time = time.time()\n",
249
- "\n",
250
- " for i in tqdm(range(current_index, len(df)), desc=\"Annotating\"):\n",
251
- " row = df.iloc[i]\n",
252
- "\n",
253
- " # Query BLOOMZ (equivalent to DeepSeek query)\n",
254
- " response = query_bloomz_local(row['prompt'])\n",
255
- " parsed_data = parse_response(response)\n",
256
- "\n",
257
- " # Log the complete interaction for debugging\n",
258
- " log_response(\n",
259
- " idx=i,\n",
260
- " name=row.get('real_name', row.get('name', 'Unknown')),\n",
261
- " prompt=row['prompt'],\n",
262
- " full_prompt=query_bloomz_local.last_full_prompt if hasattr(query_bloomz_local, 'last_full_prompt') else 'N/A',\n",
263
- " raw_response=response if response else 'None',\n",
264
- " parsed_data=parsed_data\n",
265
- " )\n",
266
- "\n",
267
- " # Update dataframe\n",
268
- " for key, value in parsed_data.items():\n",
269
- " df.at[i, key] = value\n",
270
- "\n",
271
- " current_index = i + 1\n",
272
- "\n",
273
- " # Save progress at intervals\n",
274
- " if (i + 1) % SAVE_INTERVAL == 0 or (i + 1) == len(df):\n",
275
- " df.to_csv(output_file, index=False)\n",
276
- " with open(index_file, 'w') as f:\n",
277
- " f.write(str(current_index))\n",
278
- " print(f\"✅ Progress saved after {i+1} rows\")\n",
279
- "\n",
280
- " # Optional: Add small delay to prevent overheating (not needed for rate limiting like DeepSeek)\n",
281
- " # time.sleep(0.1)\n",
282
- "\n",
283
- " elapsed_total = time.time() - start_time\n",
284
- " print(f\"\\n✅ Done! Final results saved to {output_file}\")\n",
285
- "\n",
286
- " # Summary statistics (same as DeepSeek)\n",
287
- " print(\"\\n=== Summary Statistics ===\")\n",
288
- " print(f\"Total processed: {len(df)}\")\n",
289
- " print(f\"\\nGender distribution:\")\n",
290
- " print(df['gender'].value_counts())\n",
291
- " print(f\"\\nTop 10 profession combinations:\")\n",
292
- " print(df['profession_llm'].value_counts().head(10))\n",
293
- " print(f\"\\nTop 10 countries:\")\n",
294
- " print(df['country'].value_counts().head(10))\n",
295
- "\n",
296
- " # Sample results\n",
297
- " print(\"\\n=== Sample Results ===\")\n",
298
- " display_cols = ['real_name', 'full_name', 'gender', 'profession_llm', 'country']\n",
299
- " available_cols = [col for col in display_cols if col in df.columns]\n",
300
- " print(df[available_cols].head(10).to_string(index=False))\n",
301
- "\n",
302
- " # Additional info for local model\n",
303
- " print(f\"\\nTotal time: {elapsed_total/60:.1f} minutes\")\n",
304
- " print(f\"Average speed: {len(df)/(elapsed_total/3600):.1f} samples/hour\")\n",
305
- " if torch.cuda.is_available():\n",
306
- " print(f\"Final VRAM usage: {torch.cuda.max_memory_allocated() / 1024**3:.2f} GB\")\n",
307
- "\n",
308
- "except Exception as e:\n",
309
- " print(f\"⚠️ Error encountered: {e}\")\n",
310
- " print(f\"⚠️ Last processed index: {current_index}\")\n",
311
- "\n",
312
- " # Save progress before exiting\n",
313
- " df.to_csv(output_file, index=False)\n",
314
- " with open(index_file, 'w') as f:\n",
315
- " f.write(str(current_index))\n",
316
- "\n",
317
- " print(f\"⚠️ Progress saved up to row {current_index}\")\n"
318
- ]
319
- },
320
- {
321
- "cell_type": "code",
322
- "execution_count": null,
323
- "id": "c458d50f-0cb3-421e-99ab-62795f88242c",
324
- "metadata": {},
325
- "outputs": [],
326
- "source": []
327
- }
328
- ],
329
- "metadata": {
330
- "kernelspec": {
331
- "display_name": "pm-paper",
332
- "language": "python",
333
- "name": "pm-paper"
334
- },
335
- "language_info": {
336
- "codemirror_mode": {
337
- "name": "ipython",
338
- "version": 3
339
- },
340
- "file_extension": ".py",
341
- "mimetype": "text/x-python",
342
- "name": "python",
343
- "nbconvert_exporter": "python",
344
- "pygments_lexer": "ipython3",
345
- "version": "3.11.13"
346
- }
347
- },
348
- "nbformat": 4,
349
- "nbformat_minor": 5
350
- }
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
jupyter_notebooks/Section_2-3-4_Figure_8_Step_1_LLM_annotation.ipynb ADDED
@@ -0,0 +1,509 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
+ "metadata": {},
14
+ "source": [
15
+ "### Unified Model Loading & Inference\nAbstracted interface for loading and querying Mistral, Gemma, and Qwen models."
16
+ ]
17
+ },
18
+ {
19
+ "cell_type": "code",
20
+ "metadata": {},
21
+ "source": [
22
+ "import pandas as pd\n",
23
+ "import json\n",
24
+ "import time\n",
25
+ "import re\n",
26
+ "from pathlib import Path\n",
27
+ "from tqdm import tqdm\n",
28
+ "import torch\n",
29
+ "from transformers import AutoModelForCausalLM, AutoTokenizer, BitsAndBytesConfig\n",
30
+ "import signal\n",
31
+ "from contextlib import contextmanager\n",
32
+ "\n",
33
+ "# Configuration\n",
34
+ "current_dir = Path.cwd()\n",
35
+ "CACHE_DIR = current_dir.parent / \"data/models\"\n",
36
+ "CACHE_DIR.mkdir(parents=True, exist_ok=True)\n",
37
+ "\n",
38
+ "# Model configurations\n",
39
+ "MODEL_CONFIGS = {\n",
40
+ " 'mistral': {\n",
41
+ " 'name': 'mistralai/Mistral-7B-Instruct-v0.3',\n",
42
+ " 'dtype': torch.bfloat16,\n",
43
+ " 'quantization': None,\n",
44
+ " 'generation_params': {\n",
45
+ " 'max_new_tokens': 512,\n",
46
+ " 'temperature': 0.05,\n",
47
+ " 'do_sample': True,\n",
48
+ " 'top_p': 0.8,\n",
49
+ " }\n",
50
+ " },\n",
51
+ " 'gemma': {\n",
52
+ " 'name': 'google/gemma-3-27b-it',\n",
53
+ " 'dtype': torch.bfloat16,\n",
54
+ " 'quantization': None,\n",
55
+ " 'generation_params': {\n",
56
+ " 'max_new_tokens': 512,\n",
57
+ " 'temperature': 0.1,\n",
58
+ " 'do_sample': True,\n",
59
+ " 'top_p': 0.9,\n",
60
+ " }\n",
61
+ " },\n",
62
+ " 'qwen': {\n",
63
+ " 'name': 'Qwen/Qwen2.5-32B-Instruct',\n",
64
+ " 'dtype': None, # Will use quantization\n",
65
+ " 'quantization': BitsAndBytesConfig(\n",
66
+ " load_in_8bit=True,\n",
67
+ " llm_int8_threshold=6.0,\n",
68
+ " llm_int8_has_fp16_weight=False\n",
69
+ " ),\n",
70
+ " 'generation_params': {\n",
71
+ " 'max_new_tokens': 100,\n",
72
+ " 'temperature': 0.1,\n",
73
+ " 'do_sample': False,\n",
74
+ " }\n",
75
+ " }\n",
76
+ "}\n",
77
+ "\n",
78
+ "PROFESSION_CATEGORIES = [\n",
79
+ " \"actor\",\n",
80
+ " \"adult performer\",\n",
81
+ " \"singer/musician\",\n",
82
+ " \"model\",\n",
83
+ " \"online personality\",\n",
84
+ " \"public figure\",\n",
85
+ " \"voice actor/ASMR\",\n",
86
+ " \"sports professional\",\n",
87
+ " \"tv personality\"\n",
88
+ "]\n"
89
+ ],
90
+ "execution_count": null,
91
+ "outputs": []
92
+ },
93
+ {
94
+ "cell_type": "code",
95
+ "metadata": {},
96
+ "source": [
97
+ "def load_model(model_type='mistral'):\n",
98
+ " \"\"\"\n",
99
+ " Load model and tokenizer based on type.\n",
100
+ " \n",
101
+ " Args:\n",
102
+ " model_type: 'mistral', 'gemma', or 'qwen'\n",
103
+ " \n",
104
+ " Returns:\n",
105
+ " tuple: (model, tokenizer, config)\n",
106
+ " \"\"\"\n",
107
+ " if model_type not in MODEL_CONFIGS:\n",
108
+ " raise ValueError(f\"Unknown model type: {model_type}. Choose from {list(MODEL_CONFIGS.keys())}\")\n",
109
+ " \n",
110
+ " config = MODEL_CONFIGS[model_type]\n",
111
+ " model_name = config['name']\n",
112
+ " \n",
113
+ " device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n",
114
+ " print(f\"Loading model: {model_name}\")\n",
115
+ " print(f\"Cache directory: {CACHE_DIR}\")\n",
116
+ " print(f\"Device: {device}\\n\")\n",
117
+ " \n",
118
+ " if device == \"cpu\":\n",
119
+ " print(\"⚠️ WARNING: No GPU detected! Inference will be VERY slow.\")\n",
120
+ " \n",
121
+ " # Load tokenizer\n",
122
+ " try:\n",
123
+ " tokenizer = AutoTokenizer.from_pretrained(\n",
124
+ " model_name,\n",
125
+ " cache_dir=str(CACHE_DIR),\n",
126
+ " use_fast=True\n",
127
+ " )\n",
128
+ " except:\n",
129
+ " tokenizer = AutoTokenizer.from_pretrained(\n",
130
+ " model_name,\n",
131
+ " cache_dir=str(CACHE_DIR),\n",
132
+ " use_fast=False\n",
133
+ " )\n",
134
+ " \n",
135
+ " if tokenizer.pad_token is None:\n",
136
+ " tokenizer.pad_token = tokenizer.eos_token\n",
137
+ " \n",
138
+ " # Load model\n",
139
+ " model_kwargs = {\n",
140
+ " 'cache_dir': str(CACHE_DIR),\n",
141
+ " 'device_map': 'auto',\n",
142
+ " 'trust_remote_code': False\n",
143
+ " }\n",
144
+ " \n",
145
+ " if config['quantization']:\n",
146
+ " model_kwargs['quantization_config'] = config['quantization']\n",
147
+ " else:\n",
148
+ " model_kwargs['torch_dtype'] = config['dtype']\n",
149
+ " \n",
150
+ " model = AutoModelForCausalLM.from_pretrained(model_name, **model_kwargs)\n",
151
+ " model.eval()\n",
152
+ " \n",
153
+ " # Check VRAM\n",
154
+ " if torch.cuda.is_available():\n",
155
+ " vram_gb = torch.cuda.max_memory_allocated() / 1024**3\n",
156
+ " print(f\"VRAM used: {vram_gb:.2f} GB\\n\")\n",
157
+ " \n",
158
+ " return model, tokenizer, config\n"
159
+ ],
160
+ "execution_count": null,
161
+ "outputs": []
162
+ },
163
+ {
164
+ "cell_type": "code",
165
+ "metadata": {},
166
+ "source": [
167
+ "@contextmanager\n",
168
+ "def timeout(duration):\n",
169
+ " \"\"\"Context manager for timeout.\"\"\"\n",
170
+ " def handler(signum, frame):\n",
171
+ " raise TimeoutError(\"Operation timed out\")\n",
172
+ " \n",
173
+ " signal.signal(signal.SIGALRM, handler)\n",
174
+ " signal.alarm(duration)\n",
175
+ " try:\n",
176
+ " yield\n",
177
+ " finally:\n",
178
+ " signal.alarm(0)\n",
179
+ "\n",
180
+ "def query_model(prompt, model, tokenizer, config, use_timeout=False):\n",
181
+ " \"\"\"\n",
182
+ " Query model with given prompt.\n",
183
+ " \n",
184
+ " Args:\n",
185
+ " prompt: Input prompt string\n",
186
+ " model: Loaded model\n",
187
+ " tokenizer: Loaded tokenizer\n",
188
+ " config: Model configuration dict\n",
189
+ " use_timeout: Whether to use 60s timeout (for Qwen)\n",
190
+ " \n",
191
+ " Returns:\n",
192
+ " str: Model response or None on error\n",
193
+ " \"\"\"\n",
194
+ " try:\n",
195
+ " device = next(model.parameters()).device\n",
196
+ " \n",
197
+ " # Format as chat message\n",
198
+ " messages = [\n",
199
+ " {\"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",
200
+ " {\"role\": \"user\", \"content\": prompt}\n",
201
+ " ]\n",
202
+ " \n",
203
+ " # Tokenize\n",
204
+ " if hasattr(tokenizer, 'apply_chat_template'):\n",
205
+ " text = tokenizer.apply_chat_template(\n",
206
+ " messages,\n",
207
+ " tokenize=False,\n",
208
+ " add_generation_prompt=True\n",
209
+ " )\n",
210
+ " else:\n",
211
+ " text = f\"[INST] {prompt} [/INST]\"\n",
212
+ " \n",
213
+ " inputs = tokenizer([text], return_tensors=\"pt\", padding=True).to(device)\n",
214
+ " \n",
215
+ " # Generation parameters\n",
216
+ " gen_kwargs = config['generation_params'].copy()\n",
217
+ " gen_kwargs['pad_token_id'] = tokenizer.eos_token_id\n",
218
+ " \n",
219
+ " # Generate\n",
220
+ " generation_fn = lambda: model.generate(**inputs, **gen_kwargs)\n",
221
+ " \n",
222
+ " if use_timeout:\n",
223
+ " with timeout(60):\n",
224
+ " with torch.no_grad():\n",
225
+ " outputs = generation_fn()\n",
226
+ " else:\n",
227
+ " with torch.no_grad():\n",
228
+ " outputs = generation_fn()\n",
229
+ " \n",
230
+ " # Decode\n",
231
+ " generated_ids = outputs[0][inputs['input_ids'].shape[1]:]\n",
232
+ " response = tokenizer.decode(generated_ids, skip_special_tokens=True)\n",
233
+ " \n",
234
+ " return response.strip()\n",
235
+ " \n",
236
+ " except TimeoutError:\n",
237
+ " print(f\"[ERROR] Generation timed out after 60 seconds\")\n",
238
+ " return None\n",
239
+ " except Exception as e:\n",
240
+ " print(f\"[ERROR] Generation failed: {e}\")\n",
241
+ " return None\n"
242
+ ],
243
+ "execution_count": null,
244
+ "outputs": []
245
+ },
246
+ {
247
+ "cell_type": "code",
248
+ "metadata": {},
249
+ "source": [
250
+ "def create_prompt(row):\n",
251
+ " \"\"\"Create annotation prompt from row data.\"\"\"\n",
252
+ " name = row['real_name'] if pd.notna(row.get('real_name')) else row.get('name', '')\n",
253
+ " \n",
254
+ " # Gather hints\n",
255
+ " hints = []\n",
256
+ " if pd.notna(row.get('likely_profession')):\n",
257
+ " hints.append(str(row['likely_profession']))\n",
258
+ " if pd.notna(row.get('likely_nationality')):\n",
259
+ " hints.append(str(row['likely_nationality']))\n",
260
+ " if pd.notna(row.get('likely_country')):\n",
261
+ " hints.append(str(row['likely_country']))\n",
262
+ " \n",
263
+ " # Add tags if needed\n",
264
+ " if len(hints) < 3:\n",
265
+ " for i in range(1, 8):\n",
266
+ " tag_col = f'tag_{i}'\n",
267
+ " if tag_col in row and pd.notna(row[tag_col]):\n",
268
+ " tag_val = str(row[tag_col])\n",
269
+ " if tag_val not in hints:\n",
270
+ " hints.append(tag_val)\n",
271
+ " if len(hints) >= 5:\n",
272
+ " break\n",
273
+ " \n",
274
+ " hint_text = \", \".join(hints[:5]) if hints else \"none\"\n",
275
+ " \n",
276
+ " return f\"\"\"Extract information about '{name}' ({hint_text}).\n",
277
+ "\n",
278
+ "Respond with EXACTLY 5 numbered lines. Each line must contain ONLY the value, no labels or extra text.\n",
279
+ "\n",
280
+ "FORMAT REQUIREMENTS:\n",
281
+ "1. Full legal name in Western order (first last). VALUE ONLY.\n",
282
+ "2. Stage names/aliases, comma-separated. If none, write \"None\". VALUE ONLY.\n",
283
+ "3. Gender: MUST be exactly one word: Male, Female, Other, or Unknown. VALUE ONLY.\n",
284
+ "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",
285
+ "5. Primary country: Country name only (e.g., \"China\", \"United States\", \"Colombia\"). VALUE ONLY.\n",
286
+ "\n",
287
+ "RULES:\n",
288
+ "- Professions MUST match the exact categories listed (actress = actor)\n",
289
+ "- \"online personality\" includes streamers, cosplayers, YouTubers, influencers\n",
290
+ "- \"public figure\" includes politicians, activists, journalists, authors\n",
291
+ "- Use \"Unknown\" when uncertain or for fictional characters\n",
292
+ "- NO explanatory text, NO labels like \"Gender:\", NO prefixes\n",
293
+ "- For multi-role people, list up to 3 categories by relevance\n",
294
+ "\n",
295
+ "EXAMPLE FORMAT:\n",
296
+ "1. Taylor Swift\n",
297
+ "2. None\n",
298
+ "3. Female\n",
299
+ "4. singer/musician, public figure\n",
300
+ "5. United States\"\"\"\n"
301
+ ],
302
+ "execution_count": null,
303
+ "outputs": []
304
+ },
305
+ {
306
+ "cell_type": "code",
307
+ "metadata": {},
308
+ "source": [
309
+ "def parse_response(response):\n",
310
+ " \"\"\"Parse model response into structured fields.\"\"\"\n",
311
+ " if not response:\n",
312
+ " return {\n",
313
+ " 'full_name': 'Unknown',\n",
314
+ " 'aliases': 'Unknown',\n",
315
+ " 'gender': 'Unknown',\n",
316
+ " 'profession_llm': 'Unknown',\n",
317
+ " 'country': 'Unknown'\n",
318
+ " }\n",
319
+ " \n",
320
+ " lines = [line.strip() for line in response.split('\\n') if line.strip()]\n",
321
+ " \n",
322
+ " fields = {\n",
323
+ " 'full_name': 'Unknown',\n",
324
+ " 'aliases': 'Unknown',\n",
325
+ " 'gender': 'Unknown',\n",
326
+ " 'profession_llm': 'Unknown',\n",
327
+ " 'country': 'Unknown'\n",
328
+ " }\n",
329
+ " \n",
330
+ " for line in lines:\n",
331
+ " if line.startswith('1.'):\n",
332
+ " fields['full_name'] = line[2:].strip()\n",
333
+ " elif line.startswith('2.'):\n",
334
+ " fields['aliases'] = line[2:].strip()\n",
335
+ " elif line.startswith('3.'):\n",
336
+ " gender_raw = line[2:].strip()\n",
337
+ " gender_raw = re.sub(r'^(Gender:|gender:)\\s*', '', gender_raw, flags=re.IGNORECASE)\n",
338
+ " gender_match = re.search(r'\\b(Male|Female|Other|Unknown)\\b', gender_raw, re.IGNORECASE)\n",
339
+ " fields['gender'] = gender_match.group(1).capitalize() if gender_match else gender_raw\n",
340
+ " elif line.startswith('4.'):\n",
341
+ " fields['profession_llm'] = line[2:].strip()\n",
342
+ " elif line.startswith('5.'):\n",
343
+ " country_raw = line[2:].strip()\n",
344
+ " country_raw = re.sub(r'^(Primary\\s+)?(associated\\s+)?country:\\s*', '', country_raw, flags=re.IGNORECASE)\n",
345
+ " fields['country'] = country_raw\n",
346
+ " \n",
347
+ " return fields\n"
348
+ ],
349
+ "execution_count": null,
350
+ "outputs": []
351
+ },
352
+ {
353
+ "cell_type": "code",
354
+ "metadata": {},
355
+ "source": [
356
+ "def annotate_dataset(model_type='mistral', test_mode=False, test_size=100, max_rows=50862, save_interval=10):\n",
357
+ " \"\"\"\n",
358
+ " Annotate dataset using specified model.\n",
359
+ " \n",
360
+ " Args:\n",
361
+ " model_type: 'mistral', 'gemma', or 'qwen'\n",
362
+ " test_mode: If True, only process test_size rows\n",
363
+ " test_size: Number of rows to process in test mode\n",
364
+ " max_rows: Maximum rows to process\n",
365
+ " save_interval: Save progress every N rows\n",
366
+ " \"\"\"\n",
367
+ " # Setup paths\n",
368
+ " input_file = current_dir.parent / \"data/CSV/model_adapter/real_person_adapter_step_02_NER.csv\"\n",
369
+ " output_file = current_dir.parent / f\"data/CSV/{model_type}_local_annotated_POI{'_test' if test_mode else ''}.csv\"\n",
370
+ " index_file = current_dir.parent / f\"misc/query_indicies/{model_type}_local_query_index.txt\"\n",
371
+ " index_file.parent.mkdir(parents=True, exist_ok=True)\n",
372
+ " \n",
373
+ " # Load model\n",
374
+ " model, tokenizer, config = load_model(model_type)\n",
375
+ " \n",
376
+ " # Load data\n",
377
+ " print(f\"Loaded {len(df)} rows from input file\")\n",
378
+ " df = pd.read_csv(input_file)\n",
379
+ " \n",
380
+ " # Merge existing annotations if available\n",
381
+ " if output_file.exists():\n",
382
+ " existing_df = pd.read_csv(output_file)\n",
383
+ " annotation_cols = ['full_name', 'aliases', 'gender', 'profession_llm', 'country']\n",
384
+ " for col in annotation_cols:\n",
385
+ " if col in existing_df.columns:\n",
386
+ " df[col] = existing_df[col][:len(df)]\n",
387
+ " \n",
388
+ " # Apply limits\n",
389
+ " if test_mode:\n",
390
+ " df = df.head(test_size).copy()\n",
391
+ " elif max_rows:\n",
392
+ " df = df.head(max_rows).copy()\n",
393
+ " \n",
394
+ " # Create prompts\n",
395
+ " df['prompt'] = df.apply(create_prompt, axis=1)\n",
396
+ " \n",
397
+ " # Load progress index\n",
398
+ " current_index = 0\n",
399
+ " if index_file.exists():\n",
400
+ " try:\n",
401
+ " current_index = int(index_file.read_text().strip())\n",
402
+ " except:\n",
403
+ " current_index = 0\n",
404
+ " \n",
405
+ " print(f\"Resuming from index {current_index}\")\n",
406
+ " \n",
407
+ " # Process rows\n",
408
+ " use_timeout = (model_type == 'qwen')\n",
409
+ " \n",
410
+ " for i in tqdm(range(current_index, len(df)), desc=f\"{model_type.capitalize()} Annotation\"):\n",
411
+ " prompt = df.at[i, \"prompt\"]\n",
412
+ " \n",
413
+ " # Query with retries\n",
414
+ " response = None\n",
415
+ " for attempt in range(3):\n",
416
+ " response = query_model(prompt, model, tokenizer, config, use_timeout)\n",
417
+ " \n",
418
+ " if response and len(response.strip()) > 10:\n",
419
+ " break\n",
420
+ " \n",
421
+ " print(f\"⚠️ Row {i}: Empty or invalid response, retry {attempt+1}/3\")\n",
422
+ " time.sleep(0.5)\n",
423
+ " \n",
424
+ " # Skip if invalid\n",
425
+ " if not response or len(response.strip()) <= 10:\n",
426
+ " print(f\"❌ Row {i}: failed after retries, skipping\")\n",
427
+ " continue\n",
428
+ " \n",
429
+ " # Parse and validate\n",
430
+ " parsed = parse_response(response)\n",
431
+ " \n",
432
+ " if all(v == \"Unknown\" for v in parsed.values()):\n",
433
+ " print(f\"❌ Row {i}: parsed as all Unknown, skipping\")\n",
434
+ " continue\n",
435
+ " \n",
436
+ " # Write fields\n",
437
+ " for key, value in parsed.items():\n",
438
+ " df.at[i, key] = value\n",
439
+ " \n",
440
+ " current_index = i + 1\n",
441
+ " \n",
442
+ " # GPU cleanup\n",
443
+ " if torch.cuda.is_available():\n",
444
+ " torch.cuda.empty_cache()\n",
445
+ " torch.cuda.synchronize()\n",
446
+ " \n",
447
+ " # Save progress\n",
448
+ " if (i + 1) % save_interval == 0 or (i + 1) == len(df):\n",
449
+ " df.to_csv(output_file, index=False)\n",
450
+ " index_file.write_text(str(current_index))\n",
451
+ " print(f\"💾 Progress saved after row {i+1}\")\n",
452
+ " \n",
453
+ " # Final save\n",
454
+ " df.to_csv(output_file, index=False)\n",
455
+ " index_file.write_text(str(current_index))\n",
456
+ " print(f\"✓ Finished annotation with {model_type}\")\n"
457
+ ],
458
+ "execution_count": null,
459
+ "outputs": []
460
+ },
461
+ {
462
+ "cell_type": "markdown",
463
+ "metadata": {},
464
+ "source": [
465
+ "### Usage Examples\nRun annotation with your chosen model."
466
+ ]
467
+ },
468
+ {
469
+ "cell_type": "code",
470
+ "metadata": {},
471
+ "source": [
472
+ "# Example 1: Annotate with Mistral (13.5 GB VRAM)\n",
473
+ "# annotate_dataset(model_type='mistral', test_mode=False)\n",
474
+ "\n",
475
+ "# Example 2: Annotate with Gemma (56.3 GB VRAM)\n",
476
+ "# annotate_dataset(model_type='gemma', test_mode=False)\n",
477
+ "\n",
478
+ "# Example 3: Annotate with Qwen (32.7 GB VRAM, 8-bit)\n",
479
+ "# annotate_dataset(model_type='qwen', test_mode=False)\n",
480
+ "\n",
481
+ "# Test mode (first 100 rows)\n",
482
+ "# annotate_dataset(model_type='mistral', test_mode=True, test_size=100)\n"
483
+ ],
484
+ "execution_count": null,
485
+ "outputs": []
486
+ }
487
+ ],
488
+ "metadata": {
489
+ "kernelspec": {
490
+ "display_name": "latm",
491
+ "language": "python",
492
+ "name": "python3"
493
+ },
494
+ "language_info": {
495
+ "codemirror_mode": {
496
+ "name": "ipython",
497
+ "version": 3
498
+ },
499
+ "file_extension": ".py",
500
+ "mimetype": "text/x-python",
501
+ "name": "python",
502
+ "nbconvert_exporter": "python",
503
+ "pygments_lexer": "ipython3",
504
+ "version": "3.10.15"
505
+ }
506
+ },
507
+ "nbformat": 4,
508
+ "nbformat_minor": 5
509
+ }
jupyter_notebooks/{Section_2-3-4_compare-models.ipynb → Section_2-3-4_Figure_8_Step_2_response_comparison_and_consensus_extraction.ipynb} RENAMED
File without changes
jupyter_notebooks/Section_2-3-4_Figure_8_deepfake_adapters-Copy1.ipynb DELETED
The diff for this file is too large to render. See raw diff
 
jupyter_notebooks/Section_2-3-4_Figure_8_deepfake_adapters.ipynb DELETED
The diff for this file is too large to render. See raw diff
 
jupyter_notebooks/Section_2-3-4__Figure_8_Deepfake_victims.ipynb DELETED
@@ -1,724 +0,0 @@
1
- {
2
- "cells": [
3
- {
4
- "cell_type": "markdown",
5
- "id": "06763fde",
6
- "metadata": {},
7
- "source": [
8
- "# LLM annotation of Deepfake adapters"
9
- ]
10
- },
11
- {
12
- "cell_type": "markdown",
13
- "id": "b773045a",
14
- "metadata": {},
15
- "source": [
16
- "## Step 01 Data cleaning NER using spaCy"
17
- ]
18
- },
19
- {
20
- "cell_type": "markdown",
21
- "id": "234a55e5",
22
- "metadata": {},
23
- "source": [
24
- "#### Here we clean leetspeak and architectre specifiers from the names as a preprossesing step for the named entity recognition NER below"
25
- ]
26
- },
27
- {
28
- "cell_type": "code",
29
- "execution_count": 21,
30
- "id": "f177df11",
31
- "metadata": {
32
- "execution": {
33
- "iopub.execute_input": "2025-11-21T10:32:27.999485Z",
34
- "iopub.status.busy": "2025-11-21T10:32:27.999298Z",
35
- "iopub.status.idle": "2025-11-21T10:32:28.001776Z",
36
- "shell.execute_reply": "2025-11-21T10:32:28.001287Z",
37
- "shell.execute_reply.started": "2025-11-21T10:32:27.999469Z"
38
- }
39
- },
40
- "outputs": [],
41
- "source": [
42
- "import pandas as pd\n",
43
- "import spacy\n",
44
- "import re\n",
45
- "import torch\n",
46
- "from pathlib import Path\n",
47
- "import unicodedata\n"
48
- ]
49
- },
50
- {
51
- "cell_type": "code",
52
- "execution_count": 22,
53
- "id": "a2383f32",
54
- "metadata": {
55
- "execution": {
56
- "iopub.execute_input": "2025-11-21T10:32:28.388500Z",
57
- "iopub.status.busy": "2025-11-21T10:32:28.388376Z",
58
- "iopub.status.idle": "2025-11-21T10:32:30.176781Z",
59
- "shell.execute_reply": "2025-11-21T10:32:30.176262Z",
60
- "shell.execute_reply.started": "2025-11-21T10:32:28.388488Z"
61
- }
62
- },
63
- "outputs": [],
64
- "source": [
65
- "current_dir = Path.cwd()\n",
66
- "poi_models_dir = current_dir.parent / \"data/CSV/model_adapter/real_person_adapters.csv\" ### POI models dataset\n",
67
- "output = current_dir.parent / \"data/CSV/model_adapter/real_person_adapter_step_01_NER.csv\" ### Output file"
68
- ]
69
- },
70
- {
71
- "cell_type": "code",
72
- "execution_count": 15,
73
- "id": "66fc691f",
74
- "metadata": {
75
- "execution": {
76
- "iopub.execute_input": "2025-11-21T10:23:43.615765Z",
77
- "iopub.status.busy": "2025-11-21T10:23:43.615607Z",
78
- "iopub.status.idle": "2025-11-21T10:23:47.279107Z",
79
- "shell.execute_reply": "2025-11-21T10:23:47.278503Z",
80
- "shell.execute_reply.started": "2025-11-21T10:23:43.615751Z"
81
- }
82
- },
83
- "outputs": [],
84
- "source": [
85
- "nlp = spacy.load(\"en_core_web_sm\") # or another model of your choice"
86
- ]
87
- },
88
- {
89
- "cell_type": "code",
90
- "execution_count": 16,
91
- "id": "8348c4a4",
92
- "metadata": {
93
- "execution": {
94
- "iopub.execute_input": "2025-11-21T10:23:47.279657Z",
95
- "iopub.status.busy": "2025-11-21T10:23:47.279527Z",
96
- "iopub.status.idle": "2025-11-21T10:24:34.043631Z",
97
- "shell.execute_reply": "2025-11-21T10:24:34.042977Z",
98
- "shell.execute_reply.started": "2025-11-21T10:23:47.279644Z"
99
- }
100
- },
101
- "outputs": [
102
- {
103
- "name": "stdout",
104
- "output_type": "stream",
105
- "text": [
106
- "Done! Saved to NER_poi_step_01.csv\n"
107
- ]
108
- }
109
- ],
110
- "source": [
111
- "def preprocess_name(name):\n",
112
- " name = str(name)\n",
113
- "\n",
114
- " # Normalize unicode characters (e.g., fancy fonts)\n",
115
- " name = unicodedata.normalize(\"NFKD\", name)\n",
116
- "\n",
117
- " # Lowercase everything\n",
118
- " name = name.lower()\n",
119
- "\n",
120
- " # Remove special keywords and patterns\n",
121
- " junk_words = [\n",
122
- " 'jav', 'jp', 'lora', 'locon', 'lycoris', 'requested', 'japanese', 'model',\n",
123
- " 'flux', 'flux1.d', 'pony', 'realistic'\n",
124
- " ]\n",
125
- " for word in junk_words:\n",
126
- " name = re.sub(rf'\\b{re.escape(word)}\\b', '', name, flags=re.IGNORECASE)\n",
127
- "\n",
128
- " # Remove versions like v1, v2.0, etc.\n",
129
- " name = re.sub(r'v\\.?\\d+(\\.\\d+)?', '', name)\n",
130
- "\n",
131
- " # Remove 'not' followed by a word\n",
132
- " name = re.sub(r'\\bnot\\s+\\w+', '', name)\n",
133
- "\n",
134
- " # Replace underscores and pipes with spaces\n",
135
- " name = re.sub(r'[_|]', ' ', name)\n",
136
- "\n",
137
- " # Remove parentheses and content within\n",
138
- " name = re.sub(r'\\(.*?\\)', '', name)\n",
139
- "\n",
140
- " # Remove excess whitespace\n",
141
- " name = re.sub(r'\\s+', ' ', name).strip()\n",
142
- "\n",
143
- " return name\n",
144
- "\n",
145
- "# -------------------------\n",
146
- "# Fallback Extractor\n",
147
- "# -------------------------\n",
148
- "def fallback_extract(text):\n",
149
- " words = text.split()\n",
150
- " capitalized = [w for w in words if w and w[0].isalpha()]\n",
151
- " if len(capitalized) >= 2:\n",
152
- " return \" \".join(capitalized[:2])\n",
153
- " elif capitalized:\n",
154
- " return capitalized[0]\n",
155
- " return None\n",
156
- "\n",
157
- "# -------------------------\n",
158
- "# Full Extraction Logic\n",
159
- "# -------------------------\n",
160
- "def extract_real_name(raw_name):\n",
161
- " cleaned = preprocess_name(raw_name)\n",
162
- " doc = nlp(cleaned)\n",
163
- " persons = [ent.text for ent in doc.ents if ent.label_ == \"PERSON\"]\n",
164
- " if persons:\n",
165
- " return persons[0]\n",
166
- " return fallback_extract(cleaned)\n",
167
- "\n",
168
- "# -------------------------\n",
169
- "# Load Data and Apply\n",
170
- "# -------------------------\n",
171
- "df = pd.read_csv(poi_models_dir) # Or your own path\n",
172
- "\n",
173
- "# Apply the full extractor\n",
174
- "texts = df['name'].astype(str).tolist()\n",
175
- "docs = list(nlp.pipe([preprocess_name(t) for t in texts], batch_size=32))\n",
176
- "\n",
177
- "def extract_from_doc(doc, raw_text):\n",
178
- " persons = [ent.text for ent in doc.ents if ent.label_ == \"PERSON\"]\n",
179
- " if persons:\n",
180
- " return persons[0]\n",
181
- " return fallback_extract(preprocess_name(raw_text))\n",
182
- "\n",
183
- "df['real_name'] = [extract_from_doc(doc, raw_text) for doc, raw_text in zip(docs, texts)]\n",
184
- "\n",
185
- "# Save the result\n",
186
- "df.to_csv(output, index=False)\n",
187
- "print(\"Done! Saved to NER_poi_step_01.csv\")\n"
188
- ]
189
- },
190
- {
191
- "cell_type": "markdown",
192
- "id": "a40b17fe",
193
- "metadata": {},
194
- "source": [
195
- "## Step 2: compare country and profession with lists"
196
- ]
197
- },
198
- {
199
- "cell_type": "code",
200
- "execution_count": 17,
201
- "id": "414954ed",
202
- "metadata": {
203
- "execution": {
204
- "iopub.execute_input": "2025-11-21T10:27:41.992098Z",
205
- "iopub.status.busy": "2025-11-21T10:27:41.991889Z",
206
- "iopub.status.idle": "2025-11-21T10:27:49.779590Z",
207
- "shell.execute_reply": "2025-11-21T10:27:49.778889Z",
208
- "shell.execute_reply.started": "2025-11-21T10:27:41.992081Z"
209
- }
210
- },
211
- "outputs": [
212
- {
213
- "name": "stdout",
214
- "output_type": "stream",
215
- "text": [
216
- " real_name likely_country likely_nationality likely_profession\n",
217
- "0 iu South Korean celebrity\n",
218
- "1 super pose \n",
219
- "2 liyuu \n",
220
- "3 irene South Korean celebrity\n",
221
- "4 aespa karina South Korean celebrity\n"
222
- ]
223
- }
224
- ],
225
- "source": [
226
- "import pandas as pd\n",
227
- "from pathlib import Path\n",
228
- "\n",
229
- "# Set up paths\n",
230
- "current_dir = Path.cwd()\n",
231
- "countries = current_dir.parent / \"misc/lists/countries.csv\"\n",
232
- "professions = current_dir.parent / \"misc/lists/professions.csv\"\n",
233
- "inputNER = current_dir.parent / \"data/CSV/model_adapter/real_person_adapter_step_01_NER.csv\"\n",
234
- "outfile = current_dir.parent / \"data/CSV/model_adapter/real_person_adapter_step_02_NER.csv\"\n",
235
- "\n",
236
- "# Load datasets\n",
237
- "poi_df = pd.read_csv(inputNER)\n",
238
- "countries_df = pd.read_csv(countries)\n",
239
- "professions_df = pd.read_csv(professions)\n",
240
- "\n",
241
- "# Step 1: Combine tags into one lowercase list\n",
242
- "def combine_tags(row):\n",
243
- " return [str(row[f\"tag_{i}\"]).strip().lower() for i in range(1, 8) if pd.notna(row.get(f\"tag_{i}\"))]\n",
244
- "\n",
245
- "poi_df[\"tags\"] = poi_df.apply(combine_tags, axis=1)\n",
246
- "\n",
247
- "# Step 2: Build tag → (country, nationality) mapping\n",
248
- "tag_to_country_nationality = {}\n",
249
- "\n",
250
- "for _, row in countries_df.iterrows():\n",
251
- " country = str(row[\"en_short_name\"]).strip()\n",
252
- " nationality = str(row[\"nationality\"]).strip()\n",
253
- "\n",
254
- " country_lc = country.lower()\n",
255
- " nationality_lc = nationality.lower()\n",
256
- "\n",
257
- " # Add variations to the mapping\n",
258
- " tag_to_country_nationality[country_lc] = (country, \"\")\n",
259
- " tag_to_country_nationality[nationality_lc] = (\"\", nationality)\n",
260
- " tag_to_country_nationality[country_lc.replace(\" \", \"\")] = (country, \"\")\n",
261
- " tag_to_country_nationality[nationality_lc.replace(\" \", \"\")] = (\"\", nationality)\n",
262
- "\n",
263
- " for part in country_lc.split():\n",
264
- " tag_to_country_nationality[part] = (country, \"\")\n",
265
- " for part in nationality_lc.split():\n",
266
- " tag_to_country_nationality[part] = (\"\", nationality)\n",
267
- "\n",
268
- "# Step 3: Infer likely_country and likely_nationality\n",
269
- "# Step 3: Infer likely_country and likely_nationality\n",
270
- "def infer_country_and_nationality(tags):\n",
271
- " for tag in tags:\n",
272
- " cleaned = tag.replace(\" \", \"\").lower()\n",
273
- " if cleaned in tag_to_country_nationality:\n",
274
- " country, nationality = tag_to_country_nationality[cleaned]\n",
275
- " # Special case: skip if country is \"Isle of Man\"\n",
276
- " if country == \"Isle of Man\":\n",
277
- " country = \"\"\n",
278
- " return pd.Series([country, nationality])\n",
279
- " return pd.Series([\"\", \"\"])\n",
280
- "\n",
281
- "\n",
282
- "poi_df[[\"likely_country\", \"likely_nationality\"]] = poi_df[\"tags\"].apply(infer_country_and_nationality)\n",
283
- "\n",
284
- "# Step 4: Build tag → profession mapping\n",
285
- "profession_alias_map = {}\n",
286
- "\n",
287
- "for _, row in professions_df.iterrows():\n",
288
- " canonical = str(row['profession']).strip().lower()\n",
289
- " profession_alias_map[canonical] = canonical\n",
290
- " for alias_col in ['alias_1', 'alias_2', 'alias_3']:\n",
291
- " alias = row.get(alias_col)\n",
292
- " if pd.notna(alias):\n",
293
- " profession_alias_map[str(alias).strip().lower()] = canonical\n",
294
- "\n",
295
- "# Step 5: Infer likely profession from tags\n",
296
- "def infer_profession_from_tags(tags):\n",
297
- " matched = []\n",
298
- " for tag in tags:\n",
299
- " cleaned = tag.strip().lower()\n",
300
- " if cleaned in profession_alias_map:\n",
301
- " matched.append(profession_alias_map[cleaned])\n",
302
- "\n",
303
- " if not matched:\n",
304
- " return \"\"\n",
305
- " if \"celebrity\" in matched and len(set(matched)) > 1:\n",
306
- " # Drop 'celebrity' if other professions are present\n",
307
- " matched = [m for m in matched if m != \"celebrity\"]\n",
308
- "\n",
309
- " return matched[0] # Return the first specific match\n",
310
- "\n",
311
- "\n",
312
- "poi_df[\"likely_profession\"] = poi_df[\"tags\"].apply(infer_profession_from_tags)\n",
313
- "\n",
314
- "# Step 6: Save enriched dataset\n",
315
- "poi_df.to_csv(outfile, index=False)\n",
316
- "\n",
317
- "# Optional: Preview\n",
318
- "print(poi_df[[\"real_name\", \"likely_country\", \"likely_nationality\", \"likely_profession\"]].head())\n"
319
- ]
320
- },
321
- {
322
- "cell_type": "code",
323
- "execution_count": 18,
324
- "id": "054f230b",
325
- "metadata": {
326
- "execution": {
327
- "iopub.execute_input": "2025-11-21T10:27:58.250798Z",
328
- "iopub.status.busy": "2025-11-21T10:27:58.250588Z",
329
- "iopub.status.idle": "2025-11-21T10:27:58.253063Z",
330
- "shell.execute_reply": "2025-11-21T10:27:58.252570Z",
331
- "shell.execute_reply.started": "2025-11-21T10:27:58.250780Z"
332
- }
333
- },
334
- "outputs": [],
335
- "source": [
336
- "#!pip install transformers torch\n",
337
- "#!python -m spacy download en_core_web_trf\n",
338
- "#pip install openai"
339
- ]
340
- },
341
- {
342
- "cell_type": "markdown",
343
- "id": "59185461",
344
- "metadata": {},
345
- "source": [
346
- "## Step 3: Query Deepseek-v3 with NAME and HINTS"
347
- ]
348
- },
349
- {
350
- "cell_type": "code",
351
- "execution_count": 3,
352
- "id": "504b970f",
353
- "metadata": {
354
- "execution": {
355
- "iopub.execute_input": "2025-11-21T10:02:17.101296Z",
356
- "iopub.status.busy": "2025-11-21T10:02:17.101174Z",
357
- "iopub.status.idle": "2025-11-21T10:02:17.103118Z",
358
- "shell.execute_reply": "2025-11-21T10:02:17.102724Z",
359
- "shell.execute_reply.started": "2025-11-21T10:02:17.101283Z"
360
- }
361
- },
362
- "outputs": [],
363
- "source": [
364
- "#!pip install openpyxl"
365
- ]
366
- },
367
- {
368
- "cell_type": "code",
369
- "execution_count": 19,
370
- "id": "7c209115",
371
- "metadata": {
372
- "execution": {
373
- "iopub.execute_input": "2025-11-21T10:29:28.176565Z",
374
- "iopub.status.busy": "2025-11-21T10:29:28.176373Z",
375
- "iopub.status.idle": "2025-11-21T10:30:01.404034Z",
376
- "shell.execute_reply": "2025-11-21T10:30:01.403269Z",
377
- "shell.execute_reply.started": "2025-11-21T10:29:28.176550Z"
378
- }
379
- },
380
- "outputs": [
381
- {
382
- "name": "stdout",
383
- "output_type": "stream",
384
- "text": [
385
- "Row 1/4...\n"
386
- ]
387
- },
388
- {
389
- "ename": "ModuleNotFoundError",
390
- "evalue": "No module named 'openpyxl'",
391
- "output_type": "error",
392
- "traceback": [
393
- "\u001b[31m---------------------------------------------------------------------------\u001b[39m",
394
- "\u001b[31mModuleNotFoundError\u001b[39m Traceback (most recent call last)",
395
- "\u001b[36mCell\u001b[39m\u001b[36m \u001b[39m\u001b[32mIn[19]\u001b[39m\u001b[32m, line 112\u001b[39m\n\u001b[32m 110\u001b[39m out_df.to_csv(OUTPUT_CSV, index=\u001b[38;5;28;01mFalse\u001b[39;00m)\n\u001b[32m 111\u001b[39m \u001b[38;5;66;03m# Excel\u001b[39;00m\n\u001b[32m--> \u001b[39m\u001b[32m112\u001b[39m \u001b[43mout_df\u001b[49m\u001b[43m.\u001b[49m\u001b[43mto_excel\u001b[49m\u001b[43m(\u001b[49m\u001b[43mOUTPUT_XLSX\u001b[49m\u001b[43m,\u001b[49m\u001b[43m \u001b[49m\u001b[43mindex\u001b[49m\u001b[43m=\u001b[49m\u001b[38;5;28;43;01mFalse\u001b[39;49;00m\u001b[43m)\u001b[49m\n\u001b[32m 113\u001b[39m \u001b[38;5;28;01mwith\u001b[39;00m \u001b[38;5;28mopen\u001b[39m(INDEX_FILE, \u001b[33m'\u001b[39m\u001b[33mw\u001b[39m\u001b[33m'\u001b[39m) \u001b[38;5;28;01mas\u001b[39;00m f:\n\u001b[32m 114\u001b[39m f.write(\u001b[38;5;28mstr\u001b[39m(current_index))\n",
396
- "\u001b[36mFile \u001b[39m\u001b[32m/shares/weddigen.ki.uzh/laura_wagner/phase_01/pm-paper/.venv/lib/python3.11/site-packages/pandas/util/_decorators.py:333\u001b[39m, in \u001b[36mdeprecate_nonkeyword_arguments.<locals>.decorate.<locals>.wrapper\u001b[39m\u001b[34m(*args, **kwargs)\u001b[39m\n\u001b[32m 327\u001b[39m \u001b[38;5;28;01mif\u001b[39;00m \u001b[38;5;28mlen\u001b[39m(args) > num_allow_args:\n\u001b[32m 328\u001b[39m warnings.warn(\n\u001b[32m 329\u001b[39m msg.format(arguments=_format_argument_list(allow_args)),\n\u001b[32m 330\u001b[39m \u001b[38;5;167;01mFutureWarning\u001b[39;00m,\n\u001b[32m 331\u001b[39m stacklevel=find_stack_level(),\n\u001b[32m 332\u001b[39m )\n\u001b[32m--> \u001b[39m\u001b[32m333\u001b[39m \u001b[38;5;28;01mreturn\u001b[39;00m \u001b[43mfunc\u001b[49m\u001b[43m(\u001b[49m\u001b[43m*\u001b[49m\u001b[43margs\u001b[49m\u001b[43m,\u001b[49m\u001b[43m \u001b[49m\u001b[43m*\u001b[49m\u001b[43m*\u001b[49m\u001b[43mkwargs\u001b[49m\u001b[43m)\u001b[49m\n",
397
- "\u001b[36mFile \u001b[39m\u001b[32m/shares/weddigen.ki.uzh/laura_wagner/phase_01/pm-paper/.venv/lib/python3.11/site-packages/pandas/core/generic.py:2439\u001b[39m, in \u001b[36mNDFrame.to_excel\u001b[39m\u001b[34m(self, excel_writer, sheet_name, na_rep, float_format, columns, header, index, index_label, startrow, startcol, engine, merge_cells, inf_rep, freeze_panes, storage_options, engine_kwargs)\u001b[39m\n\u001b[32m 2426\u001b[39m \u001b[38;5;28;01mfrom\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[34;01mpandas\u001b[39;00m\u001b[34;01m.\u001b[39;00m\u001b[34;01mio\u001b[39;00m\u001b[34;01m.\u001b[39;00m\u001b[34;01mformats\u001b[39;00m\u001b[34;01m.\u001b[39;00m\u001b[34;01mexcel\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[38;5;28;01mimport\u001b[39;00m ExcelFormatter\n\u001b[32m 2428\u001b[39m formatter = ExcelFormatter(\n\u001b[32m 2429\u001b[39m df,\n\u001b[32m 2430\u001b[39m na_rep=na_rep,\n\u001b[32m (...)\u001b[39m\u001b[32m 2437\u001b[39m inf_rep=inf_rep,\n\u001b[32m 2438\u001b[39m )\n\u001b[32m-> \u001b[39m\u001b[32m2439\u001b[39m \u001b[43mformatter\u001b[49m\u001b[43m.\u001b[49m\u001b[43mwrite\u001b[49m\u001b[43m(\u001b[49m\n\u001b[32m 2440\u001b[39m \u001b[43m \u001b[49m\u001b[43mexcel_writer\u001b[49m\u001b[43m,\u001b[49m\n\u001b[32m 2441\u001b[39m \u001b[43m \u001b[49m\u001b[43msheet_name\u001b[49m\u001b[43m=\u001b[49m\u001b[43msheet_name\u001b[49m\u001b[43m,\u001b[49m\n\u001b[32m 2442\u001b[39m \u001b[43m \u001b[49m\u001b[43mstartrow\u001b[49m\u001b[43m=\u001b[49m\u001b[43mstartrow\u001b[49m\u001b[43m,\u001b[49m\n\u001b[32m 2443\u001b[39m \u001b[43m \u001b[49m\u001b[43mstartcol\u001b[49m\u001b[43m=\u001b[49m\u001b[43mstartcol\u001b[49m\u001b[43m,\u001b[49m\n\u001b[32m 2444\u001b[39m \u001b[43m \u001b[49m\u001b[43mfreeze_panes\u001b[49m\u001b[43m=\u001b[49m\u001b[43mfreeze_panes\u001b[49m\u001b[43m,\u001b[49m\n\u001b[32m 2445\u001b[39m \u001b[43m \u001b[49m\u001b[43mengine\u001b[49m\u001b[43m=\u001b[49m\u001b[43mengine\u001b[49m\u001b[43m,\u001b[49m\n\u001b[32m 2446\u001b[39m \u001b[43m \u001b[49m\u001b[43mstorage_options\u001b[49m\u001b[43m=\u001b[49m\u001b[43mstorage_options\u001b[49m\u001b[43m,\u001b[49m\n\u001b[32m 2447\u001b[39m \u001b[43m \u001b[49m\u001b[43mengine_kwargs\u001b[49m\u001b[43m=\u001b[49m\u001b[43mengine_kwargs\u001b[49m\u001b[43m,\u001b[49m\n\u001b[32m 2448\u001b[39m \u001b[43m\u001b[49m\u001b[43m)\u001b[49m\n",
398
- "\u001b[36mFile \u001b[39m\u001b[32m/shares/weddigen.ki.uzh/laura_wagner/phase_01/pm-paper/.venv/lib/python3.11/site-packages/pandas/io/formats/excel.py:943\u001b[39m, in \u001b[36mExcelFormatter.write\u001b[39m\u001b[34m(self, writer, sheet_name, startrow, startcol, freeze_panes, engine, storage_options, engine_kwargs)\u001b[39m\n\u001b[32m 941\u001b[39m need_save = \u001b[38;5;28;01mFalse\u001b[39;00m\n\u001b[32m 942\u001b[39m \u001b[38;5;28;01melse\u001b[39;00m:\n\u001b[32m--> \u001b[39m\u001b[32m943\u001b[39m writer = \u001b[43mExcelWriter\u001b[49m\u001b[43m(\u001b[49m\n\u001b[32m 944\u001b[39m \u001b[43m \u001b[49m\u001b[43mwriter\u001b[49m\u001b[43m,\u001b[49m\n\u001b[32m 945\u001b[39m \u001b[43m \u001b[49m\u001b[43mengine\u001b[49m\u001b[43m=\u001b[49m\u001b[43mengine\u001b[49m\u001b[43m,\u001b[49m\n\u001b[32m 946\u001b[39m \u001b[43m \u001b[49m\u001b[43mstorage_options\u001b[49m\u001b[43m=\u001b[49m\u001b[43mstorage_options\u001b[49m\u001b[43m,\u001b[49m\n\u001b[32m 947\u001b[39m \u001b[43m \u001b[49m\u001b[43mengine_kwargs\u001b[49m\u001b[43m=\u001b[49m\u001b[43mengine_kwargs\u001b[49m\u001b[43m,\u001b[49m\n\u001b[32m 948\u001b[39m \u001b[43m \u001b[49m\u001b[43m)\u001b[49m\n\u001b[32m 949\u001b[39m need_save = \u001b[38;5;28;01mTrue\u001b[39;00m\n\u001b[32m 951\u001b[39m \u001b[38;5;28;01mtry\u001b[39;00m:\n",
399
- "\u001b[36mFile \u001b[39m\u001b[32m/shares/weddigen.ki.uzh/laura_wagner/phase_01/pm-paper/.venv/lib/python3.11/site-packages/pandas/io/excel/_openpyxl.py:57\u001b[39m, in \u001b[36mOpenpyxlWriter.__init__\u001b[39m\u001b[34m(self, path, engine, date_format, datetime_format, mode, storage_options, if_sheet_exists, engine_kwargs, **kwargs)\u001b[39m\n\u001b[32m 44\u001b[39m \u001b[38;5;28;01mdef\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[34m__init__\u001b[39m(\n\u001b[32m 45\u001b[39m \u001b[38;5;28mself\u001b[39m,\n\u001b[32m 46\u001b[39m path: FilePath | WriteExcelBuffer | ExcelWriter,\n\u001b[32m (...)\u001b[39m\u001b[32m 55\u001b[39m ) -> \u001b[38;5;28;01mNone\u001b[39;00m:\n\u001b[32m 56\u001b[39m \u001b[38;5;66;03m# Use the openpyxl module as the Excel writer.\u001b[39;00m\n\u001b[32m---> \u001b[39m\u001b[32m57\u001b[39m \u001b[38;5;28;01mfrom\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[34;01mopenpyxl\u001b[39;00m\u001b[34;01m.\u001b[39;00m\u001b[34;01mworkbook\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[38;5;28;01mimport\u001b[39;00m Workbook\n\u001b[32m 59\u001b[39m engine_kwargs = combine_kwargs(engine_kwargs, kwargs)\n\u001b[32m 61\u001b[39m \u001b[38;5;28msuper\u001b[39m().\u001b[34m__init__\u001b[39m(\n\u001b[32m 62\u001b[39m path,\n\u001b[32m 63\u001b[39m mode=mode,\n\u001b[32m (...)\u001b[39m\u001b[32m 66\u001b[39m engine_kwargs=engine_kwargs,\n\u001b[32m 67\u001b[39m )\n",
400
- "\u001b[31mModuleNotFoundError\u001b[39m: No module named 'openpyxl'"
401
- ]
402
- }
403
- ],
404
- "source": [
405
- "import pandas as pd\n",
406
- "import openai\n",
407
- "import time\n",
408
- "import os\n",
409
- "from pathlib import Path\n",
410
- "from openai import OpenAI # Add this import\n",
411
- "\n",
412
- "# === PATHS & CONFIG ===\n",
413
- "current_dir = Path.cwd()\n",
414
- "inputCSV = current_dir.parent / \"data/CSV/model_adapter/real_person_adapter_step_02_NER.csv\"\n",
415
- "api_key_file = current_dir.parent / \"misc/credentials/deepseek_api_key.txt\" #store your API key under misc/credentials/deepseek_api_key.txt\n",
416
- "\n",
417
- "# Output both CSV and Excel for compatibility\n",
418
- "OUTPUT_CSV = current_dir.parent / \"data/CSV/model_adapter/real_person_adapter_step_03_deepseek.csv\"\n",
419
- "OUTPUT_XLSX = current_dir.parent / \"data/CSV/model_adapter/real_person_adapter_step_03_deepseek.xlsx\"\n",
420
- "INDEX_FILE = current_dir.parent / \"misc/query-indicies/deepseek_query_index.txt\"\n",
421
- "SAVE_INTERVAL = 1 # Save every N rows\n",
422
- "START_ROW = 1 # Row index to start from (0-based)\n",
423
- "END_ROW = 5 # Row index to end (exclusive)\n",
424
- "\n",
425
- "# === LOAD API KEY & CLIENT ===\n",
426
- "with open(api_key_file) as f:\n",
427
- " api_key = f.read().strip()\n",
428
- "\n",
429
- "client = OpenAI(\n",
430
- " api_key=api_key,\n",
431
- " base_url=\"https://api.deepseek.com/v1\"\n",
432
- ")\n",
433
- "\n",
434
- "# === LOAD DATA ===\n",
435
- "df = pd.read_csv(inputCSV)\n",
436
- "df = df.iloc[START_ROW:END_ROW].reset_index(drop=True)\n",
437
- "\n",
438
- "# === PREPARE PROMPTS ===\n",
439
- "def create_prompt(row):\n",
440
- " name = row['real_name'] if pd.notna(row['real_name']) else row['name']\n",
441
- " hints = []\n",
442
- " for col in ('likely_profession','likely_nationality','likely_country'):\n",
443
- " if pd.notna(row.get(col, None)):\n",
444
- " hints.append(row[col])\n",
445
- " if not hints:\n",
446
- " tags = [row[f'tag_{i}'] for i in range(1,8) if pd.notna(row.get(f'tag_{i}', None))]\n",
447
- " hints.extend(tags[:3])\n",
448
- " hint_text = \", \".join(hints[:5])\n",
449
- " return (\n",
450
- " f\"Given '{name}' ({hint_text}), provide:\\n\"\n",
451
- " \"1. Full legal name (Western order if non-latin script)\\n\"\n",
452
- " \"2. Any stage names/aliases (comma separated)\\n\"\n",
453
- " \"3. Gender\\n\"\n",
454
- " \"4. Top 3 most specific, factual professions (use industry-standard terms, no euphemisms)\\n\"\n",
455
- " \"5. Primary country associated\\n\"\n",
456
- " \"Use 'Unknown' when uncertain or you encounter a fictional character or place. \"\n",
457
- " \"For entertainment fields, specify sub-genres when known (kpop, adult industry, etc.).\"\n",
458
- " )\n",
459
- "\n",
460
- "prompts = df.apply(create_prompt, axis=1).tolist()\n",
461
- "\n",
462
- "# === CHECK FOR EXISTING OUTPUT ===\n",
463
- "if os.path.exists(INDEX_FILE):\n",
464
- " with open(INDEX_FILE, 'r') as f:\n",
465
- " current_index = int(f.read().strip())\n",
466
- "else:\n",
467
- " current_index = 0\n",
468
- "\n",
469
- "results = []\n",
470
- "if os.path.exists(OUTPUT_XLSX):\n",
471
- " existing = pd.read_excel(OUTPUT_XLSX)\n",
472
- " if {'full_name','aliases','gender','profession_llm','country'}.issubset(existing.columns):\n",
473
- " results = existing[['full_name','aliases','gender','profession_llm','country']].values.tolist()\n",
474
- "\n",
475
- "# === QUERY & PARSE ===\n",
476
- "def query_deepseek(prompt):\n",
477
- " try:\n",
478
- " resp = client.chat.completions.create(\n",
479
- " model=\"deepseek-chat\",\n",
480
- " messages=[\n",
481
- " {\"role\":\"system\",\"content\":\"You extract key data on a person; respond with exactly 5 numbered lines.\"},\n",
482
- " {\"role\":\"user\",\"content\":prompt}\n",
483
- " ],\n",
484
- " temperature=0.05, top_p=0.8\n",
485
- " )\n",
486
- " return resp.choices[0].message.content.strip()\n",
487
- " except Exception as e:\n",
488
- " print(\"API error:\", e)\n",
489
- " return \"\"\n",
490
- "\n",
491
- "def parse_response(resp):\n",
492
- " lines = [l.strip() for l in resp.split('\\n') if l.strip()]\n",
493
- " out = [\"Unknown\"]*5\n",
494
- " for l in lines:\n",
495
- " if l.startswith('1.'): out[0] = l[2:].strip()\n",
496
- " elif l.startswith('2.'): out[1] = l[2:].strip()\n",
497
- " elif l.startswith('3.'): out[2] = l[2:].strip()\n",
498
- " elif l.startswith('4.'): out[3] = l[2:].strip()\n",
499
- " elif l.startswith('5.'): out[4] = l[2:].strip()\n",
500
- " return out\n",
501
- "\n",
502
- "# === PROCESS & SAVE ===\n",
503
- "for i in range(current_index, len(df)):\n",
504
- " print(f\"Row {i+1}/{len(df)}...\")\n",
505
- " data = parse_response(query_deepseek(prompts[i]))\n",
506
- " if i < len(results): results[i] = data\n",
507
- " else: results.append(data)\n",
508
- " current_index = i+1\n",
509
- "\n",
510
- " if current_index % SAVE_INTERVAL == 0 or current_index == len(df):\n",
511
- " out_df = df.iloc[:current_index].copy()\n",
512
- " out_df[['full_name','aliases','gender','profession_llm','country']] = pd.DataFrame(results[:current_index])\n",
513
- " # CSV\n",
514
- " out_df.to_csv(OUTPUT_CSV, index=False)\n",
515
- " # Excel\n",
516
- " out_df.to_excel(OUTPUT_XLSX, index=False)\n",
517
- " with open(INDEX_FILE, 'w') as f:\n",
518
- " f.write(str(current_index))\n",
519
- " print(\"Saved up to row\", current_index)\n",
520
- " time.sleep(1)\n",
521
- "\n",
522
- "print(\"All done! Files:\", OUTPUT_CSV, OUTPUT_XLSX)\n"
523
- ]
524
- },
525
- {
526
- "cell_type": "markdown",
527
- "id": "9d377005",
528
- "metadata": {},
529
- "source": [
530
- "# Aggregate by individual names"
531
- ]
532
- },
533
- {
534
- "cell_type": "markdown",
535
- "id": "d2e75354",
536
- "metadata": {},
537
- "source": [
538
- " ##### e.g. Emma Watson [model1, model2, model3] etc."
539
- ]
540
- },
541
- {
542
- "cell_type": "code",
543
- "execution_count": 11,
544
- "id": "747c3a2f",
545
- "metadata": {},
546
- "outputs": [],
547
- "source": [
548
- "import pandas as pd\n",
549
- "import re\n",
550
- "from pathlib import Path\n",
551
- "current_dir = Path.cwd()\n",
552
- "\n",
553
- "profession_map = current_dir.parent / \"misc/lists/mapped_professions.csv\"\n",
554
- "\n",
555
- "poi_df = current_dir.parent / \"data/CSV/Deepseek_annotated_POI.csv\"\n",
556
- "\n",
557
- "output = current_dir.parent / \"data/CSV/Deepseek_annotated_POI_aggregated.csv\"\n",
558
- "\n",
559
- "countries_csv = current_dir.parent / \"misc/lists/countries.csv\"\n",
560
- "countries_df = pd.read_csv(countries_csv)\n",
561
- "\n",
562
- "# Extract valid country names (strip whitespace)\n",
563
- "valid_countries = set(countries_df['en_short_name'].str.strip())\n",
564
- "\n",
565
- "\n",
566
- "# Load the dataset\n",
567
- "df = pd.read_csv(poi_df) # Update path if needed\n",
568
- "\n",
569
- "# Step 1: Group by 'full_name' and aggregate required information\n",
570
- "grouped_df = df.groupby('full_name').agg(\n",
571
- " No_of_models=('id', 'count'),\n",
572
- " modelIDs=('id', lambda x: list(x)),\n",
573
- " combinedDownloadCount=('downloadCount', 'sum')\n",
574
- ").reset_index()\n",
575
- "\n",
576
- "# Step 2: Keep representative info for each person\n",
577
- "# Keep representative info (including aliases)\n",
578
- "additional_columns = df.groupby('full_name').agg(\n",
579
- " country=('country', 'first'),\n",
580
- " profession_llm=('profession_llm', 'first'),\n",
581
- " gender=('gender', 'first'),\n",
582
- " aliases=('aliases', 'first') # ✅ Add this line\n",
583
- ").reset_index()\n",
584
- "\n",
585
- "\n",
586
- "\n",
587
- "def standardize_country(country):\n",
588
- " if not isinstance(country, str):\n",
589
- " return \"Unknown\"\n",
590
- "\n",
591
- " country_clean = country.strip()\n",
592
- " lowered = country_clean.lower()\n",
593
- "\n",
594
- " # Handle fictional or fantasy countries\n",
595
- " fictional_keywords = [\"fictional\", \"westeros\", \"asgard\", \"middle-earth\", \"naboo\", \"middle earth\", \"latveria\"]\n",
596
- " if any(keyword in lowered for keyword in fictional_keywords):\n",
597
- " return \"Unknown\"\n",
598
- "\n",
599
- " # Handle known region-based adjustments\n",
600
- " if \"macau\" in lowered:\n",
601
- " return \"Macau\"\n",
602
- " elif \"hong kong\" in lowered:\n",
603
- " return \"Hong Kong\"\n",
604
- " elif \"taiwan\" in lowered:\n",
605
- " return \"Taiwan\"\n",
606
- "\n",
607
- " # Normalize complex or alternate country names\n",
608
- " lowered = lowered.replace(\"United Kingdom of Great Britain and Northern Ireland\", \"united kingdom\")\n",
609
- " lowered = lowered.replace(\"england\", \"united kingdom\")\n",
610
- " lowered = lowered.replace(\"united states of america\", \"united states\")\n",
611
- "\n",
612
- " # Remove anything in brackets and after commas\n",
613
- " country_clean = re.sub(r\"\\(.*?\\)\", \"\", country_clean)\n",
614
- " country_clean = country_clean.split(',')[0].strip().lower()\n",
615
- "\n",
616
- " # Manual overrides\n",
617
- " replacements = {\n",
618
- " \"united kingdom\": \"UK\",\n",
619
- " \"united kingdom of great britain and northern ireland\": \"UK\",\n",
620
- " \"french southern territories\": \"Other\",\n",
621
- " \"united states\": \"US\",\n",
622
- " \"united states of america\": \"US\",\n",
623
- " \"turkey\": \"Türkiye\",\n",
624
- " \"czech republic\": \"Czechia\"\n",
625
- " }\n",
626
- "\n",
627
- " if country_clean in replacements:\n",
628
- " return replacements[country_clean]\n",
629
- "\n",
630
- " # Final check against valid country list (case-insensitive)\n",
631
- " for valid in valid_countries:\n",
632
- " if country_clean == valid.lower():\n",
633
- " return valid\n",
634
- "\n",
635
- " return \"Unknown\"\n",
636
- "\n",
637
- "\n",
638
- "# Updated function to fully remove anything in brackets (complete or not)\n",
639
- "def get_profession_short(profession):\n",
640
- " if isinstance(profession, str):\n",
641
- " # Get first part before comma\n",
642
- " first_prof = profession.split(',')[0].strip()\n",
643
- " # Remove all bracketed content, even malformed\n",
644
- " first_prof = re.sub(r\"[\\[].∗?[\\[].*?[\\]]\", \"\", first_prof) # removes properly closed\n",
645
- " first_prof = re.sub(r\"[\\(\\[].*\", \"\", first_prof) # removes malformed\n",
646
- " cleaned = first_prof.strip()\n",
647
- " # Normalize 'Actress' to 'Actor'\n",
648
- " if cleaned.lower() == \"actress\":\n",
649
- " return \"Actor\"\n",
650
- " return cleaned\n",
651
- " return None\n",
652
- "\n",
653
- "# Load your mapping file\n",
654
- "mapping_df = pd.read_csv(profession_map, on_bad_lines='skip') # or 'warn'\n",
655
- "\n",
656
- "\n",
657
- "# Ensure the mapping columns are named correctly\n",
658
- "# (Assuming columns are: 'profession_llm' or 'profession_short', and 'category' or 'mapped_profession')\n",
659
- "# Adjust these as needed\n",
660
- "mapping_df.columns = [col.strip().lower() for col in mapping_df.columns]\n",
661
- "\n",
662
- "# Rename for clarity and consistency\n",
663
- "if 'profession_llm' in mapping_df.columns:\n",
664
- " mapping_df = mapping_df.rename(columns={'profession_llm': 'profession_short'})\n",
665
- "if 'category' in mapping_df.columns:\n",
666
- " mapping_df = mapping_df.rename(columns={'category': 'mapped_profession'})\n",
667
- "\n",
668
- "# Merge the mapped profession into final_df\n",
669
- "#final_df = final_df.merge(mapping_df[['profession_short', 'mapped_profession']], on='profession_short', how='left')\n",
670
- "\n",
671
- "\n",
672
- "additional_columns = df.groupby('full_name').agg(\n",
673
- " country=('country', 'first'),\n",
674
- " profession_llm=('profession_llm', 'first'),\n",
675
- " gender=('gender', 'first'),\n",
676
- " aliases=('aliases', 'first') # <-- Added this line\n",
677
- ").reset_index()\n",
678
- "\n",
679
- "\n",
680
- "# Step 3: Merge the aggregated info with the representative info\n",
681
- "final_df = pd.merge(grouped_df, additional_columns, on='full_name', how='left')\n",
682
- "\n",
683
- "# Step 4: Clean and transform columns\n",
684
- "final_df['profession_short'] = final_df['profession_llm'].apply(get_profession_short)\n",
685
- "final_df['country'] = final_df['country'].apply(standardize_country)\n",
686
- "\n",
687
- "# Step 5: Merge with profession mapping\n",
688
- "final_df = final_df.merge(mapping_df[['profession_short', 'mapped_profession']], on='profession_short', how='left')\n",
689
- "\n",
690
- "# Optional: Save the result to a CSV file\n",
691
- "final_df.to_csv(output, index=False)\n"
692
- ]
693
- },
694
- {
695
- "cell_type": "code",
696
- "execution_count": null,
697
- "id": "704e5246",
698
- "metadata": {},
699
- "outputs": [],
700
- "source": []
701
- }
702
- ],
703
- "metadata": {
704
- "kernelspec": {
705
- "display_name": "pm-paper",
706
- "language": "python",
707
- "name": "pm-paper"
708
- },
709
- "language_info": {
710
- "codemirror_mode": {
711
- "name": "ipython",
712
- "version": 3
713
- },
714
- "file_extension": ".py",
715
- "mimetype": "text/x-python",
716
- "name": "python",
717
- "nbconvert_exporter": "python",
718
- "pygments_lexer": "ipython3",
719
- "version": "3.11.13"
720
- }
721
- },
722
- "nbformat": 4,
723
- "nbformat_minor": 5
724
- }
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
jupyter_notebooks/Section_3-3-4_deepfakes.ipynb DELETED
@@ -1,124 +0,0 @@
1
- {
2
- "cells": [
3
- {
4
- "cell_type": "markdown",
5
- "id": "a58d08a9",
6
- "metadata": {},
7
- "source": [
8
- "# LLM annotation of Deepfake adapter dataset"
9
- ]
10
- },
11
- {
12
- "cell_type": "markdown",
13
- "id": "38c3a6a6",
14
- "metadata": {},
15
- "source": [
16
- "## Step 1 NER"
17
- ]
18
- },
19
- {
20
- "cell_type": "code",
21
- "execution_count": null,
22
- "id": "ad478ed2",
23
- "metadata": {},
24
- "outputs": [],
25
- "source": [
26
- "import pandas as pd\n",
27
- "from pathlib import Path\n",
28
- "\n",
29
- "# Set up paths\n",
30
- "current_dir = Path.cwd()\n",
31
- "countries = current_dir.parent / \"misc/lists/countries.csv\"\n",
32
- "professions = current_dir.parent / \"misc/lists/professions.csv\"\n",
33
- "inputNER = current_dir.parent / \"data/CSV/NER_POI_step01_pre_annotation.csv\"\n",
34
- "outfile = current_dir.parent / \"data/CSV/NER_POI_step02_annotated.csv\"\n",
35
- "\n",
36
- "# Load datasets\n",
37
- "poi_df = pd.read_csv(inputNER)\n",
38
- "countries_df = pd.read_csv(countries)\n",
39
- "professions_df = pd.read_csv(professions)\n",
40
- "\n",
41
- "# Step 1: Combine tags into one lowercase list\n",
42
- "def combine_tags(row):\n",
43
- " return [str(row[f\"tag_{i}\"]).strip().lower() for i in range(1, 8) if pd.notna(row.get(f\"tag_{i}\"))]\n",
44
- "\n",
45
- "poi_df[\"tags\"] = poi_df.apply(combine_tags, axis=1)\n",
46
- "\n",
47
- "# Step 2: Build tag → (country, nationality) mapping\n",
48
- "tag_to_country_nationality = {}\n",
49
- "\n",
50
- "for _, row in countries_df.iterrows():\n",
51
- " country = str(row[\"en_short_name\"]).strip()\n",
52
- " nationality = str(row[\"nationality\"]).strip()\n",
53
- "\n",
54
- " country_lc = country.lower()\n",
55
- " nationality_lc = nationality.lower()\n",
56
- "\n",
57
- " # Add variations to the mapping\n",
58
- " tag_to_country_nationality[country_lc] = (country, \"\")\n",
59
- " tag_to_country_nationality[nationality_lc] = (\"\", nationality)\n",
60
- " tag_to_country_nationality[country_lc.replace(\" \", \"\")] = (country, \"\")\n",
61
- " tag_to_country_nationality[nationality_lc.replace(\" \", \"\")] = (\"\", nationality)\n",
62
- "\n",
63
- " for part in country_lc.split():\n",
64
- " tag_to_country_nationality[part] = (country, \"\")\n",
65
- " for part in nationality_lc.split():\n",
66
- " tag_to_country_nationality[part] = (\"\", nationality)\n",
67
- "\n",
68
- "# Step 3: Infer likely_country and likely_nationality\n",
69
- "def infer_country_and_nationality(tags):\n",
70
- " for tag in tags:\n",
71
- " cleaned = tag.replace(\" \", \"\").lower()\n",
72
- " if cleaned in tag_to_country_nationality:\n",
73
- " country, nationality = tag_to_country_nationality[cleaned]\n",
74
- " return pd.Series([country, nationality])\n",
75
- " return pd.Series([\"\", \"\"])\n",
76
- "\n",
77
- "poi_df[[\"likely_country\", \"likely_nationality\"]] = poi_df[\"tags\"].apply(infer_country_and_nationality)\n",
78
- "\n",
79
- "# Step 4: Build tag → profession mapping\n",
80
- "profession_alias_map = {}\n",
81
- "\n",
82
- "for _, row in professions_df.iterrows():\n",
83
- " canonical = str(row['profession']).strip().lower()\n",
84
- " profession_alias_map[canonical] = canonical\n",
85
- " for alias_col in ['alias_1', 'alias_2', 'alias_3']:\n",
86
- " alias = row.get(alias_col)\n",
87
- " if pd.notna(alias):\n",
88
- " profession_alias_map[str(alias).strip().lower()] = canonical\n",
89
- "\n",
90
- "# Step 5: Infer likely profession from tags\n",
91
- "def infer_profession_from_tags(tags):\n",
92
- " matched = []\n",
93
- " for tag in tags:\n",
94
- " cleaned = tag.strip().lower()\n",
95
- " if cleaned in profession_alias_map:\n",
96
- " matched.append(profession_alias_map[cleaned])\n",
97
- "\n",
98
- " if not matched:\n",
99
- " return \"\"\n",
100
- " if \"celebrity\" in matched and len(set(matched)) > 1:\n",
101
- " # Drop 'celebrity' if other professions are present\n",
102
- " matched = [m for m in matched if m != \"celebrity\"]\n",
103
- "\n",
104
- " return matched[0] # Return the first specific match\n",
105
- "\n",
106
- "\n",
107
- "poi_df[\"likely_profession\"] = poi_df[\"tags\"].apply(infer_profession_from_tags)\n",
108
- "\n",
109
- "# Step 6: Save enriched dataset\n",
110
- "poi_df.to_csv(outfile, index=False)\n",
111
- "\n",
112
- "# Optional: Preview\n",
113
- "print(poi_df[[\"real_name\", \"likely_country\", \"likely_nationality\", \"likely_profession\"]].head())\n"
114
- ]
115
- }
116
- ],
117
- "metadata": {
118
- "language_info": {
119
- "name": "python"
120
- }
121
- },
122
- "nbformat": 4,
123
- "nbformat_minor": 5
124
- }
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
{jupyter_notebooks → md}/DEEPFAKE_PIPELINE_GUIDE.md RENAMED
File without changes
{jupyter_notebooks → md}/LLM_MODELS_COMPARISON.md RENAMED
File without changes
{jupyter_notebooks → md}/QUICK_START_LOCAL.md RENAMED
File without changes
{jupyter_notebooks → md}/QWEN_LOCAL_SETUP.md RENAMED
File without changes
{jupyter_notebooks → md}/SPACY_NER_EXPLANATION.md RENAMED
File without changes
{jupyter_notebooks → md}/TESTING_INSTRUCTIONS.md RENAMED
File without changes
{jupyter_notebooks → md}/UPDATES_AND_FIXES.md RENAMED
File without changes