From 1fa4bb5ea47ff8372937225e86e2069007c6dea3 Mon Sep 17 00:00:00 2001 From: wassname <1103714+wassname@users.noreply.github.com> Date: Sat, 4 Oct 2025 15:24:57 +0800 Subject: [PATCH] restore --- nbs/simple_icm.py | 102 +++++++++++++++++++--------------------------- 1 file changed, 42 insertions(+), 60 deletions(-) diff --git a/nbs/simple_icm.py b/nbs/simple_icm.py index df229ce..2c2c182 100644 --- a/nbs/simple_icm.py +++ b/nbs/simple_icm.py @@ -23,9 +23,6 @@ from loguru import logger from openrouter_wrapper.logprobs import openrouter_completion_wlogprobs, get_logprobs_choices, LogprobsNotSupportedError # User's wrapper from typing import List, Tuple import asyncio -from aiocache import SimpleMemoryCache, cached, Cache - -cache = SimpleMemoryCache() try: from IPython import get_ipython @@ -55,13 +52,13 @@ class Config: num_seed: int = 8 max_iters: int = 950 # Small for demo; increase for more n_shots: int = 16 # Number of in-context examples - model_id: str = "meta-llama/llama-3.1-8b-instruct" # Logprobs supported + model: str = "meta-llama/llama-3.1-8b-instruct" # Logprobs supported provider_whitelist: Tuple[str] = None # None to let OpenRouter choose out_dir: Path = Path("../outputs/icm") # Directory to save outputs log_interval: int = 50 # Log progress every N iterations C = Config( - model_id="qwen/qwen3-235b-a22b-2507", + model="qwen/qwen3-235b-a22b-2507", provider_whitelist=[ 'Chutes','Nebius',], ) C.out_dir.mkdir(parents=True, exist_ok=True) @@ -132,16 +129,16 @@ logger.info("Initialized labels: {}", {k: v['label'] for k, v in demonstrations. def print_messages(messages): return "\n".join([f"**{m['role'].upper()}**: {m['content']}" for m in messages]) -@cached(cache) async def cached_openrouter_completion_wlogprobs(*args, **kwargs): # TODO: Async caching (e.g., with aiocache) return await openrouter_completion_wlogprobs(*args, **kwargs) -async def predict_label(example_uid, full_demos, labeled_demos, config=C, verbose=False): + +async def predict_label(example_uid, current_demos, config=C, verbose=False): # simplified from https://github.com/Jiaxin-Wen/Unsupervised-Elicitation/blob/master/src/experiments/ICM.py#L254 # Sort by consistency_key for relevance, limit to 16 relevant_demos = sorted( - [demo for uid, demo in labeled_demos.items() if uid != example_uid], + [demo for uid, demo in current_demos.items() if uid != example_uid and demo['label'] is not None], key=lambda x: x.get('consistency_key', '') )[:C.n_shots] # TODO: Use all/subsampled demos like original (group/sort by consistency_id for relevance), handling context limits (high value, med complexity). @@ -150,51 +147,49 @@ async def predict_label(example_uid, full_demos, labeled_demos, config=C, verbos instruction = "Carefully examine the noisy example pairs of (Candidate, Set). Then as consistently and coherently as possible complete the pattern, you may provide reasoning after.\n\n" # if verbose>1: - # instruction = "Examine the examples below and determine the category 0 or 1 for the final claim based on the pattern. After - - # Randomize order to avoid positional bias - random.shuffle(relevant_demos) + # instruction = "Examine the examples below and determine the category 0 or 1 for the final claim based on the pattern. After 0 or 1 output a single word to describe how category 1 compares to category 0.\n\n" + fewshot = [] + # FIXME to not lead the unsupervised model, we should avoid true/false or even 0/1 and try to use neutral labels like A/B or similar + for idx, demo in enumerate(relevant_demos): + label_str = "A" if demo['label'] == 1 else "B" + fewshot.append(f"\nCandidate: {demo['prompt']}\nSet: {label_str}\n\n") + target_prompt = demonstrations[example_uid]['prompt'] messages = [ - {"role": "system", "content": instruction} + {"role": "user", "content": instruction+"".join(fewshot)}, + {"role": "assistant", "content": f"Candidate: {target_prompt}\n"} # Assistant prefill to ensure ] + - for demo in relevant_demos: - messages.append({"role": "user", "content": demo['prompt']}) - messages.append({"role": "assistant", "content": str(demo['label'])}) - - target_demo = full_demos[example_uid] - messages.append({"role": "user", "content": target_demo['prompt']}) - + response = await cached_openrouter_completion_wlogprobs( + model_id=config.model_id, + provider_whitelist=config.provider_whitelist, + messages=messages, + max_completion_tokens=90 if verbose else 5, + temperature=0.4, + top_logprobs=8, + ) + if verbose: - logger.info(f"Predicting for UID {example_uid} with {len(relevant_demos)} shots:") - logger.info(print_messages(messages[:3])) # Log first few for brevity - if verbose > 1: - logger.info(print_messages(messages)) - + 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"--- End Debug ---") + try: - completion = await cached_openrouter_completion_wlogprobs( - model_id=config.model_id, - messages=messages, - max_completion_tokens=60 if verbose else 4, - temperature=0.0, - logprobs=True, - top_logprobs=5, - provider_whitelist=config.provider_whitelist if config.provider_whitelist else None - ) - - # Extract logprobs for '0' and '1' - choice_logp_dict, logp_dict = get_logprobs_choices(completion, choices=['A', 'B'], regex='[AB]') - - - if verbose: - logger.info(f"Predicted: {pred}, logprob: {lprob:.4f}") - - return pred, lprob + choice_logp, all_logp = get_logprobs_choices(response, ["A", "B"]) + probmass = np.exp(np.array(list(choice_logp.values()))).sum() + if probmass < 0.5: + model_response = response['choices'][0]['message']['content'] + logger.warning(f"Low prob mass {probmass:.2f} for UID {example_uid}, may indicate model confusion. Instead we got these top logprobs: {all_logp} and this output: {model_response}") + score = choice_logp["A"] - choice_logp["B"] + predicted = 1 if score > 0 else 0 + return predicted, float(score) except Exception as e: - logger.error(f"Error in predict_label for UID {example_uid}: {e}") + logger.error(f"API error: {e}") return random.choice([0, 1]), 0.0 + # %% [code] # Compute energy and metrics def compute_energy(demos, config=C): @@ -282,8 +277,8 @@ async def fix_inconsistencies_simple(demos, config=C, max_fixes=5): # Re-predict labels with new context to update scores current_labeled = {k: v for k, v in temp_demos.items() if v['label'] is not None} - new_label1, score1 = await predict_label(uid1, temp_demos, current_labeled, config) - new_label2, score2 = await predict_label(uid2, temp_demos, current_labeled, config) + new_label1, score1 = await predict_label(uid1, current_labeled, config) + new_label2, score2 = await predict_label(uid2, current_labeled, config) temp_demos[uid1]['score'] = score1 temp_demos[uid2]['score'] = score2 @@ -313,15 +308,6 @@ async def run_icm(demonstrations, config=C): current_labeled = {k: v for k, v in demonstrations.items() if v['label'] is not None} old_energy, old_metrics = compute_energy(demonstrations, config) - # Set initial scores for seeds by predicting (using current labels) - initial_labeled = {k: v for k, v in demonstrations.items() if v['label'] is not None} - for uid in initial_labeled: - pred_label, score = await predict_label(uid, demonstrations, initial_labeled, config, verbose=False) - demonstrations[uid]['score'] = score # Keep original label, update score - initial_labeled[uid]['score'] = score - old_energy, old_metrics = compute_energy(demonstrations, config) - logger.info("Initial scores set. Updated energy: {}", old_energy) - for iter in range(config.max_iters): # Weighted sampling for inconsistent groups groups = {} @@ -372,7 +358,7 @@ async def run_icm(demonstrations, config=C): verbose = 2 else: verbose = 0 - new_label, score = await predict_label(example_uid, demonstrations, current_labeled, config, verbose=verbose) + new_label, score = await predict_label(example_uid, current_labeled, config, verbose=verbose) # Update with new label and fix any inconsistencies temp_demos = demonstrations.copy() @@ -380,10 +366,6 @@ async def run_icm(demonstrations, config=C): temp_demos[example_uid]['score'] = score temp_demos = await fix_inconsistencies_simple(temp_demos, config) - # TODO check if we should enable following - # if new_label != demonstrations[example_uid].get('label', None): - # current_labeled = {k: v for k, v in temp_demos.items() if v['label'] is not None} - # Compute new energy new_energy, new_metrics = compute_energy(temp_demos, config) delta = new_energy - old_energy