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