mirror of
https://github.com/wassname/Unsupervised-Elicitation.git
synced 2026-09-09 11:16:07 +08:00
fix keyboard interrupt
This commit is contained in:
Binary file not shown.
|
Before Width: | Height: | Size: 40 KiB |
Binary file not shown.
|
Before Width: | Height: | Size: 41 KiB After Width: | Height: | Size: 33 KiB |
@@ -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",
|
||||
|
||||
@@ -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
|
||||
@@ -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
@@ -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')
|
||||
|
||||
|
||||
@@ -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" },
|
||||
|
||||
Reference in New Issue
Block a user