diff --git a/nbs/simple_icm.py b/nbs/simple_icm.py index 7ab67d6..a866eca 100644 --- a/nbs/simple_icm.py +++ b/nbs/simple_icm.py @@ -38,7 +38,7 @@ dotenv.load_dotenv() # Setup loguru logger.remove() -logger.add(sys.stderr, format="{time} | {level} | {message}", colorize=True, level="INFO") +logger.add(sys.stderr, format="{time:YYYY-MM-DD HH:mm} | {level} | {message}", colorize=True, level="DEBUG") # %% [code] @@ -51,29 +51,25 @@ class Config: beta: float = 2.0 num_seed: int = 8 max_iters: int = 950 # Small for demo; increase for more - n_shots: int = 16 # Number of in-context examples + n_shots: int = 8 # 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 out_dir: Path = Path("../outputs/icm") # Directory to save outputs - log_interval: int = 50 # Log progress every N iterations + log_interval: int = 20 # Log progress every N iterations C = Config( - model_id="qwen/qwen3-235b-a22b-2507", + model_id="qwen/qwen3-235b-a22b-2507", # $0.2 0.6 provider_whitelist=[ 'Chutes','Nebius',], ) C.out_dir.mkdir(parents=True, exist_ok=True) # C = Config( -# model_id="qwen/qwen3-30b-a3b-instruct-2507", +# model_id="qwen/qwen3-30b-a3b-instruct-2507", # 0.08 $0.33 # provider_whitelist=[ 'Chutes','Nebius',], # ) # C = Config( -# model_id="meta-llama/llama-3.1-70b-instruct", -# provider_whitelist=[ 'Cerebras','Nebius',], -# ) -# C = Config( -# model_id="meta-llama/llama-3.1-8b-instruct", +# model_id="meta-llama/llama-3.1-8b-instruct", # 0.02 0.03 # provider_whitelist=('Cerebras','Nebius',), # ) @@ -90,7 +86,7 @@ data = [] group_id = 0 groups = {} for idx, item in enumerate(dataset): - claim = f"Question: {item['Question']}\nClaim: {item['Answer']}\nI think this claim is" + claim = f"Question: {item['Question']}\nClaim: {item['Answer']}" label = item['label'] consistency_id = item['question_id'] @@ -152,12 +148,12 @@ async def predict_label(example_uid, current_demos, config=C, verbose=False): # 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") + fewshot.append(f"\nCandidate: {demo['prompt']}\nSet: {label_str}\n") target_prompt = demonstrations[example_uid]['prompt'] messages = [ - {"role": "user", "content": instruction+"".join(fewshot)}, - {"role": "assistant", "content": f"Candidate: {target_prompt}\n"} # Assistant prefill to ensure + {"role": "user", "content": instruction+"".join(fewshot)+f"Candidate: {target_prompt}\n"}, + {"role": "assistant", "content": "\n\nSet:"} # Assistant prefill to ensure ] @@ -177,16 +173,22 @@ async def predict_label(example_uid, current_demos, config=C, verbose=False): logger.info(f"--- End Debug ---") try: - choice_logp, all_logp = get_logprobs_choices(response, ["A", "B"]) - probmass = np.exp(np.array(list(choice_logp.values()))).sum() - if probmass < 0.5: + choice_strs = ["A", "B"] + choice_logp, top_logp = get_logprobs_choices(response, choice_strs, lower=False) + + + choice_in_toplogp = any([s for s in choice_strs if s in top_logp]) + + + if not choice_in_toplogp: 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}") + 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}") 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"API error: {e}") + raise e + logger.exception(f"API error: {e}") return random.choice([0, 1]), 0.0 @@ -230,7 +232,7 @@ def compute_energy(demos, config=C): energy = config.alpha * avg_lprob - num_inconsistent accuracy = np.mean([d['label'] == d['vanilla_label'] for d in labeled]) return energy, { - 'avg_prob': avg_lprob, + 'avg_lprob': avg_lprob, 'num_inconsistent': num_inconsistent, 'accuracy': accuracy, 'num_labeled': len(labeled) @@ -377,7 +379,7 @@ async def run_icm(demonstrations, config=C): demonstrations = temp_demos old_energy = new_energy current_labeled = {k: v for k, v in demonstrations.items() if v['label'] is not None} - logger.info("Iter {}: Accepted. Energy: {:.2f}. {}", iter, old_energy, accept_msg) + logger.debug("Iter {}: Accepted. Energy: {:.2f}. {}", iter, old_energy, accept_msg) else: logger.debug("Iter {}: Rejected. {}", iter, accept_msg) diff --git a/uv.lock b/uv.lock index 1bae939..d9a0948 100644 --- a/uv.lock +++ b/uv.lock @@ -26,15 +26,6 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/53/1c/8feedd607cc14c5df9aef74fe3af9a99bf660743b842a9b5b1865326b4aa/adjustText-1.3.0-py3-none-any.whl", hash = "sha256:da23d7b24b6db5ffa039bb136bfa556207365e32f48ac74b07ad26dd485bc691", size = 13154, upload-time = "2024-10-31T16:45:35.227Z" }, ] -[[package]] -name = "aiocache" -version = "0.12.3" -source = { registry = "https://pypi.org/simple" } -sdist = { url = "https://files.pythonhosted.org/packages/7a/64/b945b8025a9d1e6e2138845f4022165d3b337f55f50984fbc6a4c0a1e355/aiocache-0.12.3.tar.gz", hash = "sha256:f528b27bf4d436b497a1d0d1a8f59a542c153ab1e37c3621713cb376d44c4713", size = 132196, upload-time = "2024-09-25T13:20:23.823Z" } -wheels = [ - { url = "https://files.pythonhosted.org/packages/37/d7/15d67e05b235d1ed8c3ce61688fe4d84130e72af1657acadfaac3479f4cf/aiocache-0.12.3-py2.py3-none-any.whl", hash = "sha256:889086fc24710f431937b87ad3720a289f7fc31c4fd8b68e9f918b9bacd8270d", size = 28199, upload-time = "2024-09-25T13:20:22.688Z" }, -] - [[package]] name = "aiohappyeyeballs" version = "2.6.1" @@ -3757,7 +3748,6 @@ version = "0.1.0" source = { editable = "." } dependencies = [ { name = "adjusttext" }, - { name = "aiocache" }, { name = "alembic" }, { name = "altair" }, { name = "anthropic" }, @@ -3803,7 +3793,6 @@ dev = [ [package.metadata] requires-dist = [ { name = "adjusttext", specifier = ">=1.3.0" }, - { name = "aiocache", specifier = ">=0.12.3" }, { name = "alembic", specifier = ">=1.16.5" }, { name = "altair", specifier = ">=5.5.0" }, { name = "anthropic", specifier = ">=0.69.0" },