1- from vllm import LLM , SamplingParams
2- from rdkit import Chem
3- import re
41import pickle
2+ import re
3+
4+ import numpy as np
55import pandas as pd
6+ from rdkit import Chem
67from sklearn .model_selection import train_test_split
7- import numpy as np
8+
9+ from vllm import LLM , SamplingParams
810
911df = pd .read_csv ("/data/david/final_tasks_prompts/dataset_swapped500k_prompt.csv" )
1012
1113df = df .iloc [10_000 :].reset_index (drop = True )
1214
13- train_df , test_df = train_test_split (
14- df ,
15- test_size = 3000 ,
16- random_state = 42 ,
17- shuffle = True
18- )
15+ train_df , test_df = train_test_split (df , test_size = 3000 , random_state = 42 , shuffle = True )
1916print (test_df .columns )
2017
21- test_data = test_df .to_dict (orient = "records" )
18+ test_data = test_df .to_dict (orient = "records" )
2219
2320all_keys = set ().union (* (rec .keys () for rec in test_data ))
2421print (all_keys )
2522
2623# Load vLLM model
27- llm = LLM (model = "/data/share/sft_hf_3/" )
28- sampling_params = SamplingParams (n = 1 , presence_penalty = 0.0 , frequency_penalty = 0.0 , repetition_penalty = 1.00 , temperature = 0.8 , top_p = 0.80 , top_k = 20 , min_p = 0.0 , seed = None , stop = [], stop_token_ids = [151643 , 151644 , 151645 ], bad_words = [], include_stop_str_in_output = False , ignore_eos = False , max_tokens = 4096 , min_tokens = 0 , logprobs = None , prompt_logprobs = None , skip_special_tokens = True , spaces_between_special_tokens = True , truncate_prompt_tokens = None , guided_decoding = None )
29- #sampling_params = SamplingParams(n=5, max_tokens=4096, stop_token_ids=[151643, 151644, 151645])
24+ llm = LLM (model = "/data/share/sft_hf_3/" )
25+ sampling_params = SamplingParams (
26+ n = 1 ,
27+ presence_penalty = 0.0 ,
28+ frequency_penalty = 0.0 ,
29+ repetition_penalty = 1.00 ,
30+ temperature = 0.8 ,
31+ top_p = 0.80 ,
32+ top_k = 20 ,
33+ min_p = 0.0 ,
34+ seed = None ,
35+ stop = [],
36+ stop_token_ids = [151643 , 151644 , 151645 ],
37+ bad_words = [],
38+ include_stop_str_in_output = False ,
39+ ignore_eos = False ,
40+ max_tokens = 4096 ,
41+ min_tokens = 0 ,
42+ logprobs = None ,
43+ prompt_logprobs = None ,
44+ skip_special_tokens = True ,
45+ spaces_between_special_tokens = True ,
46+ truncate_prompt_tokens = None ,
47+ guided_decoding = None ,
48+ )
49+ # sampling_params = SamplingParams(n=5, max_tokens=4096, stop_token_ids=[151643, 151644, 151645])
50+
3051
3152def extract_answer (text ):
3253 m = re .search (r"<answer>\s*([ABCD])\s*</answer>" , text , re .IGNORECASE )
@@ -38,19 +59,20 @@ def extract_answer(text):
3859
3960 return None
4061
41- letters = ["A" ,"B" ,"C" ,"D" ]
62+
63+ letters = ["A" , "B" , "C" , "D" ]
4264
4365
44- prompts = []
45- gold_letters = []
66+ prompts = []
67+ gold_letters = []
4668for ex in test_data :
4769 true_rx = ex ["true_reaction" ]
48- fakes = [ex ["fake1" ], ex ["fake2" ], ex ["fake3" ]]
49-
70+ fakes = [ex ["fake1" ], ex ["fake2" ], ex ["fake3" ]]
71+
5072 opts = np .random .permutation ([true_rx ] + fakes ).tolist ()
51-
52- gold_letters .append ( letters [ opts .index (true_rx ) ] )
53-
73+
74+ gold_letters .append (letters [opts .index (true_rx )] )
75+
5476 prompt = (
5577 "<|im_start|>assistant\n "
5678 "You are a useful Chemistry assistant and will answer the following MCQ.\n "
@@ -79,36 +101,38 @@ def extract_answer(text):
79101for idx , (gold , out , prompt ) in enumerate (zip (gold_letters , outputs , prompts ), 1 ):
80102 print (f"\n === Example #{ idx } (gold={ gold } ) ===" )
81103 hit = False
82-
104+
83105 for j , sample in enumerate (out .outputs , 1 ):
84106 text = sample .text
85107 pred = extract_answer (text )
86- ok = ( pred == gold )
87-
108+ ok = pred == gold
109+
88110 if ok :
89111 print (f"\n *** CORRECT SAMPLE (Example { idx } , Sample { j } ) ***" )
90112 print ("Prompt:\n " , prompt )
91113 print ("Completion:\n " , text )
92114 print (f"Predicted: { pred !r} Gold: { gold !r} \n " )
93115 else :
94-
116+
95117 if np .random .rand () < 0.5 :
96118 print (f"\n *** INCORRECT SAMPLE (Example { idx } , Sample { j } ) ***" )
97119 print ("Predicted:" , pred , "Gold:" , gold )
98120 print (text [:300 ], "…\n " )
99-
121+
100122 if ok and not hit :
101123 correct += 1
102124 hit = True
103-
104- records .append ({
105- "example_idx" : idx ,
106- "gold_letter" : gold ,
107- "predicted_letter" : pred ,
108- "prompt" : prompt ,
109- "completion" : text
110- })
111- break
125+
126+ records .append (
127+ {
128+ "example_idx" : idx ,
129+ "gold_letter" : gold ,
130+ "predicted_letter" : pred ,
131+ "prompt" : prompt ,
132+ "completion" : text ,
133+ }
134+ )
135+ break
112136
113137 print (f"Example { idx } any-of-5 correct? { 'OK' if hit else 'WRONG' } " )
114138
@@ -119,10 +143,6 @@ def extract_answer(text):
119143out_df .to_csv ("/data/david/benchmark_save_completion_runs/inversion_2_correct_completions_think.csv" , index = False )
120144print (f"Saved { len (out_df )} correct completions to inversion_2_correct_completions_think.csv" )
121145
122- metrics = pd .DataFrame ([{
123- "accuracy_pct" : round (acc , 2 ),
124- "correct" : correct ,
125- "total" : len (test_data )
126- }])
146+ metrics = pd .DataFrame ([{"accuracy_pct" : round (acc , 2 ), "correct" : correct , "total" : len (test_data )}])
127147
128- metrics .to_csv ("/data/david/benchmark_save_completion_runs/inversion_2_metrics_think.csv" , index = False )
148+ metrics .to_csv ("/data/david/benchmark_save_completion_runs/inversion_2_metrics_think.csv" , index = False )
0 commit comments