fix keyboard interrupt

This commit is contained in:
wassname
2025-10-05 11:19:58 +08:00
parent 0f5c754888
commit e9f5810c08
7 changed files with 163 additions and 42 deletions
Binary file not shown.

Before

Width:  |  Height:  |  Size: 40 KiB

BIN
View File
Binary file not shown.

Before

Width:  |  Height:  |  Size: 41 KiB

After

Width:  |  Height:  |  Size: 33 KiB

+1
View File
@@ -33,6 +33,7 @@ dependencies = [
"scikit-learn>=1.7.2",
"scipy>=1.15.3",
"seaborn>=0.13.2",
"simple-parsing>=0.1.7",
"sqlalchemy==2.0.18",
"tabulate>=0.9.0",
"tenacity>=9.1.2",
-24
View File
@@ -1,24 +0,0 @@
from datasets import load_dataset
def load_daily_dilemma():
dataset = load_dataset("kellycyy/daily_dilemmas", split="test")
data = []
for item in dataset:
# Unbiased prompt - just the situation and action
prompt = f"Situation: {item['dilemma_situation']}\nAction: {item['action']}"
# Map to A/B (unrelated group names)
label = 1 if item['action_type'] == 'to_do' else 0
consistency_key = 'A' if item['action_type'] == 'to_do' else 'B'
example = {
'uid': item['idx'],
'prompt': prompt,
'vanilla_label': label,
'consistency_id': item['dilemma_idx'],
'consistency_key': consistency_key
}
data.append(example)
return data
+96
View File
@@ -0,0 +1,96 @@
from datasets import load_dataset
import pandas as pd
from collections import defaultdict
import ast
def convert_values_to_list(x):
# turn into list
s = x["values_aggregated"]
v = ast.literal_eval(s)
return {"values_aggregated": v}
def load_daily_dilemma(label_col = "label_Virtue/Truthfulness"):
dataset = load_dataset("kellycyy/daily_dilemmas", split="test")
dataset = dataset.map(convert_values_to_list)
df_labels = load_labels(dataset).rename(columns=lambda x: f"label_{x}" if x != "dilemma_idx" else x)
df = dataset.to_pandas()
df = df.merge(df_labels, on='dilemma_idx')
cols_label = [c for c in df_labels.columns if c.startswith('label_')]
# flip labels on not_to_do
df.loc[df['action_type'] == 'not_to_do', cols_label] *= -1
dataset = df.to_dict(orient='records')
data = []
for item in dataset:
# Unbiased prompt - just the situation and action
prompt = f"Situation: {item['dilemma_situation']}\nAction: {item['action']}"
# Map to A/B (unrelated group names)
consistency_key = 'A' if item['action_type'] == 'to_do' else 'B'
example = {
'uid': item['idx'],
'prompt': prompt,
'vanilla_label': item[label_col],
'consistency_id': item['dilemma_idx'],
'consistency_key': consistency_key
}
data.append(example)
return data
def load_labels(dd_dataset):
ds_values = load_dataset("kellycyy/daily_dilemmas", split="test", name="Values")
# moral tags
moral_frameworks = ["WVS", "MFT", "Virtue", "Emotion", "Maslow"]
value2framework_dicts = {}
for framework in moral_frameworks:
df_values = ds_values.to_pandas()[["value", framework]].dropna()
value2framework_dict = df_values.set_index("value")[framework].to_dict()
value2framework_dict = {k: f"{framework}/{v}" for k, v in value2framework_dict.items()}
value2framework_dicts[framework] = value2framework_dict
# make labels
df_dilemma = dd_dataset.to_pandas()[["dilemma_idx", "action_type", "values_aggregated"]]
dilemma_idx = df_dilemma["dilemma_idx"].unique()
labels = []
for d_idx in dilemma_idx:
pos_values = (
df_dilemma.query('dilemma_idx == @d_idx and action_type == "to_do"')["values_aggregated"].iloc[0].tolist()
)
neg_values = (
df_dilemma.query('dilemma_idx == @d_idx and action_type == "not_to_do"')["values_aggregated"].iloc[0].tolist()
)
label = defaultdict(int)
for framework in value2framework_dicts:
value2framework_dict = value2framework_dicts[framework]
virtues = sorted(set(value2framework_dict.values()))
pos_virtues = [value2framework_dict[k] for k in pos_values if k in value2framework_dict]
neg_virtues = [value2framework_dict[k] for k in neg_values if k in value2framework_dict]
for p in pos_virtues:
label[p] += 1
for n in neg_virtues:
label[n] -= 1
labels.append(dict(dilemma_idx=d_idx, **label))
df_labels = pd.DataFrame(labels).set_index("dilemma_idx")
assert df_labels.index.is_unique
return df_labels
+51 -18
View File
@@ -19,7 +19,7 @@ from dataclasses import dataclass, asdict
import dotenv
from loguru import logger
from openrouter_wrapper.logprobs import openrouter_completion_wlogprobs, get_logprobs_choices, LogprobsNotSupportedError # User's wrapper
from typing import List, Tuple, Callable
from typing import List, Tuple, Callable, Literal
import asyncio
from aiocache import cached
from itertools import combinations
@@ -43,6 +43,7 @@ logger.add(sys.stderr, format="<green>{time:YYYY-MM-DD HH:mm}</green> | <level>{
# Global cost tracker
total_cost = 0.0
reasoning_log = ""
# %% [code]
@@ -56,16 +57,22 @@ class Config:
num_seed: int = 8
max_iters: int = 2500 # should be at least dataset size X 2
n_shots: int = 6 # Number of in-context examples
model_id: str = "meta-llama/llama-3.1-8b-instruct" # Logprobs supported
provider_whitelist: Tuple[str] = None # None to let OpenRouter choose
# model_id: str = "meta-llama/llama-3.1-8b-instruct" # Logprobs supported
model_id: str = "qwen/qwen3-235b-a22b-2507" # Logprobs supported
provider_whitelist: Tuple[str] = ('Chutes','Nebius',) # None to let OpenRouter choose
out_dir: Path = Path("../outputs/icm") # Directory to save outputs
log_interval: int = 100 # Log progress every N iterations
dataset: Literal["truthfulqa", "daily_dilemmas"] = "truthfulqa" # Dataset name for logging
C = Config(
model_id="qwen/qwen3-235b-a22b-2507", # $0.2 0.6
provider_whitelist=[ 'Chutes','Nebius',],
)
C.out_dir.mkdir(parents=True, exist_ok=True)
import simple_parsing
C: Config = simple_parsing.parse(Config)
# C = Config(
# model_id="qwen/qwen3-235b-a22b-2507", # $0.2 0.6
# provider_whitelist=[ 'Chutes','Nebius',],
# )
# C = Config(
# model_id="qwen/qwen3-30b-a3b-instruct-2507", # 0.08 $0.33
@@ -79,14 +86,23 @@ C.out_dir.mkdir(parents=True, exist_ok=True)
logger.info(f"Config: {C}")
config_dict = asdict(C)
config_dict['out_dir'] = str(config_dict['out_dir'])
with open(C.out_dir / "icm_config.json", "w") as f:
out_dir = C.out_dir / C.dataset.replace(' ', '_').replace('/', '_')
out_dir.mkdir(parents=True, exist_ok=True)
config_dict['out_dir'] = str(C.out_dir)
with open(out_dir / "icm_config.json", "w") as f:
json.dump(config_dict, f, indent=2)
# %% [code]
from src.data.truthfulqa import load_truthfulqa, is_consistent
from data.daily_dilemmas import load_daily_dilemma
if C.dataset == "truthfulqa":
data = load_truthfulqa()
elif C.dataset == "daily_dilemmas":
data = load_daily_dilemma()
else:
raise ValueError(f"Unknown dataset {C.dataset}")
data = load_truthfulqa()
logger.info("Loaded {} examples", len(data))
# %% [code]
@@ -145,6 +161,9 @@ async def predict_label(example_uid, current_demos, config=C, verbose=False, all
{"role": "user", "content": instruction+"".join(fewshot)+f"\n\n## Candidate:\n{target_prompt}"},
{"role": "assistant", "content": "\n## Set:"} # Assistant prefill to ensure
]
if verbose>1:
messages[0]['content'] = "ALWAYS GIVE BRIEF REASONING AFTERWARDS. " + messages[0]['content']
response = await cached_openrouter_completion_wlogprobs(
@@ -152,6 +171,7 @@ async def predict_label(example_uid, current_demos, config=C, verbose=False, all
provider_whitelist=config.provider_whitelist,
messages=messages,
max_completion_tokens=160 if verbose else 5,
min_completion_tokens=30 if verbose else 1,
temperature=0.4,
top_logprobs=8,
)
@@ -160,9 +180,14 @@ async def predict_label(example_uid, current_demos, config=C, verbose=False, all
if verbose:
logger.info(f"Debug Prediction - UID {example_uid}:")
logger.info(f"messages: {print_messages(messages)}")
logger.info(f"Response Content: {response['choices'][0]['message']['content']}")
logger.info(f"messages: `{print_messages(messages)}`")
logger.info(f"Response Content: `{response['choices'][0]['message']['content']}`")
logger.info(f"--- End Debug ---")
if verbose>1:
global reasoning_log
reasoning_log += f"\n\n## Candidate:\n{target_prompt}\n## Set:\n"
reasoning_log += response['choices'][0]['message']['content']
try:
choice_strs = ["A", "B"]
@@ -174,7 +199,7 @@ async def predict_label(example_uid, current_demos, config=C, verbose=False, all
if not choice_in_toplogp:
model_response = response['choices'][0]['message']['content']
logger.warning(f"Choices not returned for UID {example_uid}, may indicate model confusion. choice_logp={choice_logp}. Instead we got these top logprobs: {top_logp} and \nmessages: {print_messages(messages)}\nthis output: {model_response}")
logger.warning(f"Choices not returned for UID {example_uid}, may indicate model confusion. choice_logp={choice_logp}. Instead we got these top logprobs: {top_logp} and \nmessages: ...`{print_messages(messages)[-90:]}`\nthis output:`{model_response}`")
score = choice_logp["A"] - choice_logp["B"]
predicted = 1 if score > 0 else 0
return predicted, float(score)
@@ -223,6 +248,9 @@ def compute_energy(demos, config=C):
energy = config.alpha * avg_lprob - num_inconsistent - (num_inconsistent / max(1, len(labeled))) # Normalized penalty
accuracy = np.mean([d['label'] == d['vanilla_label'] for d in labeled])
# flip acc if needed, as this is unsupervised
if accuracy < 0.5:
accuracy = 1 - accuracy
return energy, {
'avg_lprob': avg_lprob,
'num_inconsistent': num_inconsistent,
@@ -405,6 +433,8 @@ async def run_icm(demonstrations, config=C):
except KeyboardInterrupt:
logger.info("Stopping early.")
except asyncio.CancelledError:
logger.info("Asyncio task cancelled.")
return demonstrations, energies, accuracies
@@ -418,13 +448,13 @@ logger.info("\nFinal Results:")
logger.info("Total cost: ${:.4f}", total_cost)
logger.info("Energy: {:.2f}", final_energy)
# TODO show vanilla accuracy here for comparison
logger.info("Accuracy vs vanilla: {:.2f}, initial {:.2f}", final_metrics['accuracy'], accuracies[0])
logger.info("Accuracy [labelled] vs vanilla: {:.2f}, initial {:.2f}", final_metrics['accuracy'], accuracies[0])
logger.info("Labeled: {}/{}", final_metrics['num_labeled'], len(data))
logger.info("Inconsistencies: {}", final_metrics['num_inconsistent'])
# Final labels
df = pd.DataFrame(final_demos).T
df.to_parquet(C.out_dir / "icm_final_labels.parquet")
df.to_parquet(out_dir / "icm_final_labels.parquet")
df_labeled = df.dropna(subset='label').sort_values(by='score', key=np.abs, ascending=False)
df_labeled_disagreed = df_labeled[df_labeled['vanilla_label'] != df_labeled['label']]
@@ -435,11 +465,14 @@ print(df_labeled_disagreed[['consistency_id', 'label', 'vanilla_label', 'score',
for uid, row in df_labeled_disagreed.iterrows():
print(f"\n## Candidate: {row['prompt']}\nICM Set: {'A' if row['label']==1 else 'B'}, Vanilla Set: {'A' if row['vanilla_label']==1 else 'B'}, score={row['score']}\n")
print(f"\nFinal labeled examples saved to {out_dir / 'icm_final_labels.parquet'}")
# %% [code]
# Simple visualization (requires matplotlib)
with open(out_dir / "cost.txt", "w") as f:
f.write(f"\n\nTotal cost: ${total_cost:.4f}\n")
f.write(reasoning_log)
plt.figure(figsize=(10, 4))
plt.subplot(1, 2, 1)
@@ -450,7 +483,7 @@ plt.ylabel('Energy')
plt.subplot(1, 2, 2)
plt.plot(accuracies)
plt.title('Accuracy over Iterations')
plt.title('Accuracy [labelled] over Iterations')
plt.xlabel('Iteration')
plt.ylabel('Accuracy')
Generated
+15
View File
@@ -3452,6 +3452,19 @@ wheels = [
{ url = "https://files.pythonhosted.org/packages/e0/f9/0595336914c5619e5f28a1fb793285925a8cd4b432c9da0a987836c7f822/shellingham-1.5.4-py2.py3-none-any.whl", hash = "sha256:7ecfff8f2fd72616f7481040475a65b2bf8af90a56c89140852d1120324e8686", size = 9755, upload-time = "2023-10-24T04:13:38.866Z" },
]
[[package]]
name = "simple-parsing"
version = "0.1.7"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "docstring-parser" },
{ name = "typing-extensions" },
]
sdist = { url = "https://files.pythonhosted.org/packages/eb/c5/f1e2fcb3a81085cdf3cfed48b8c8ce0e7cc30c95dee734cbb35d6265336a/simple_parsing-0.1.7.tar.gz", hash = "sha256:225e6b35252d68f7894716101fe3bd7e6dd3d30ab7b1c3c023f77a42dbe1336f", size = 96375, upload-time = "2025-01-20T19:46:35.986Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/4f/9c/e9ea38750027a6de3e3c5e68a19fda0e7b0cd3db8045f30d0f6bc113b911/simple_parsing-0.1.7-py3-none-any.whl", hash = "sha256:5276e6c90c157362dd0173d1eecebe58361a66b457129cc9bba13b78a4e85092", size = 112782, upload-time = "2025-01-20T19:46:33.325Z" },
]
[[package]]
name = "six"
version = "1.17.0"
@@ -3782,6 +3795,7 @@ dependencies = [
{ name = "scipy", version = "1.15.3", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.11'" },
{ name = "scipy", version = "1.16.2", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.11'" },
{ name = "seaborn" },
{ name = "simple-parsing" },
{ name = "sqlalchemy" },
{ name = "tabulate" },
{ name = "tenacity" },
@@ -3827,6 +3841,7 @@ requires-dist = [
{ name = "scikit-learn", specifier = ">=1.7.2" },
{ name = "scipy", specifier = ">=1.15.3" },
{ name = "seaborn", specifier = ">=0.13.2" },
{ name = "simple-parsing", specifier = ">=0.1.7" },
{ name = "sqlalchemy", specifier = "==2.0.18" },
{ name = "tabulate", specifier = ">=0.9.0" },
{ name = "tenacity", specifier = ">=9.1.2" },