Skip to content

Commit d8e7bf6

Browse files
doncamilomclaude
andcommitted
Run black and isort for CI compliance
Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
1 parent 55c283e commit d8e7bf6

13 files changed

Lines changed: 658 additions & 693 deletions
Lines changed: 49 additions & 39 deletions
Original file line numberDiff line numberDiff line change
@@ -1,32 +1,53 @@
1-
from vllm import LLM, SamplingParams
2-
from rdkit import Chem
3-
import re
41
import pickle
2+
import re
3+
4+
import numpy as np
55
import pandas as pd
6+
from rdkit import Chem
67
from sklearn.model_selection import train_test_split
7-
import numpy as np
8+
9+
from vllm import LLM, SamplingParams
810

911
df = pd.read_csv("/data/david/final_tasks_prompts/dataset_swapped500k_prompt.csv")
1012

1113
df = 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)
1916
print(test_df.columns)
2017

21-
test_data = test_df.to_dict(orient="records")
18+
test_data = test_df.to_dict(orient="records")
2219

2320
all_keys = set().union(*(rec.keys() for rec in test_data))
2421
print(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

3152
def 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"]
4262

63+
letters = ["A", "B", "C", "D"]
4364

44-
prompts = []
45-
gold_letters = []
65+
66+
prompts = []
67+
gold_letters = []
4668
for 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"
@@ -72,10 +94,7 @@ def extract_answer(text):
7294
prompts.append(prompt)
7395

7496
if needs_retry:
75-
retry_prompts = [
76-
prompts[i] + completions[i] + "\n<answer>"
77-
for i in needs_retry
78-
]
97+
retry_prompts = [prompts[i] + completions[i] + "\n<answer>" for i in needs_retry]
7998
print(f"Retrying {len(needs_retry)} examples: {needs_retry}")
8099
retry_results = llm.generate(retry_prompts, sampling_params)
81100

@@ -98,15 +117,10 @@ def extract_answer(text):
98117
records = []
99118
for i, (gold, comp) in enumerate(zip(gold_letters, completions), start=1):
100119
pred = extract_answer(comp)
101-
hit = (pred == gold)
120+
hit = pred == gold
102121
if hit:
103122
correct += 1
104-
records.append({
105-
"example_idx": i,
106-
"gold_letter": gold,
107-
"predicted": pred,
108-
"completion": comp
109-
})
123+
records.append({"example_idx": i, "gold_letter": gold, "predicted": pred, "completion": comp})
110124
print(f"Example #{i}: gold={gold} pred={pred}{'OK' if hit else 'WRONG'}")
111125

112126
acc = 100 * correct / len(test_data)
@@ -116,10 +130,6 @@ def extract_answer(text):
116130
out_df.to_csv("/data/david/benchmark_save_completion_runs/inversion_2_correct_completions_think.csv", index=False)
117131
print(f"Saved {len(out_df)} correct completions to inversion_2_correct_completions_think.csv")
118132

119-
metrics = pd.DataFrame([{
120-
"accuracy_pct": round(acc, 2),
121-
"correct": correct,
122-
"total": len(test_data)
123-
}])
133+
metrics = pd.DataFrame([{"accuracy_pct": round(acc, 2), "correct": correct, "total": len(test_data)}])
124134

125-
metrics.to_csv("/data/david/benchmark_save_completion_runs/inversion_2_metrics_think.csv", index=False)
135+
metrics.to_csv("/data/david/benchmark_save_completion_runs/inversion_2_metrics_think.csv", index=False)
Lines changed: 62 additions & 42 deletions
Original file line numberDiff line numberDiff line change
@@ -1,32 +1,53 @@
1-
from vllm import LLM, SamplingParams
2-
from rdkit import Chem
3-
import re
41
import pickle
2+
import re
3+
4+
import numpy as np
55
import pandas as pd
6+
from rdkit import Chem
67
from sklearn.model_selection import train_test_split
7-
import numpy as np
8+
9+
from vllm import LLM, SamplingParams
810

911
df = pd.read_csv("/data/david/final_tasks_prompts/dataset_swapped500k_prompt.csv")
1012

1113
df = 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)
1916
print(test_df.columns)
2017

21-
test_data = test_df.to_dict(orient="records")
18+
test_data = test_df.to_dict(orient="records")
2219

2320
all_keys = set().union(*(rec.keys() for rec in test_data))
2421
print(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

3152
def 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 = []
4668
for 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):
79101
for 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):
119143
out_df.to_csv("/data/david/benchmark_save_completion_runs/inversion_2_correct_completions_think.csv", index=False)
120144
print(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

Comments
 (0)