Text Generation
Transformers
Safetensors
English
qwen2
chat
conversational
text-generation-inference
Instructions to use divaspoudel/iol-2026-solver with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use divaspoudel/iol-2026-solver with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("text-generation", model="divaspoudel/iol-2026-solver") messages = [ {"role": "user", "content": "Who are you?"}, ] pipe(messages)# Load model directly from transformers import AutoTokenizer, AutoModelForCausalLM tokenizer = AutoTokenizer.from_pretrained("divaspoudel/iol-2026-solver") model = AutoModelForCausalLM.from_pretrained("divaspoudel/iol-2026-solver", device_map="auto") messages = [ {"role": "user", "content": "Who are you?"}, ] inputs = tokenizer.apply_chat_template( messages, add_generation_prompt=True, tokenize=True, return_dict=True, return_tensors="pt", ).to(model.device) outputs = model.generate(**inputs, max_new_tokens=40) print(tokenizer.decode(outputs[0][inputs["input_ids"].shape[-1]:])) - Notebooks
- Google Colab
- Kaggle
- Local Apps Settings
- vLLM
How to use divaspoudel/iol-2026-solver with vLLM:
Install from pip and serve model
# Install vLLM from pip: pip install vllm # Start the vLLM server: vllm serve "divaspoudel/iol-2026-solver" # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:8000/v1/chat/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "divaspoudel/iol-2026-solver", "messages": [ { "role": "user", "content": "What is the capital of France?" } ] }'Use Docker
docker model run hf.co/divaspoudel/iol-2026-solver
- SGLang
How to use divaspoudel/iol-2026-solver with SGLang:
Install from pip and serve model
# Install SGLang from pip: pip install sglang # Start the SGLang server: python3 -m sglang.launch_server \ --model-path "divaspoudel/iol-2026-solver" \ --host 0.0.0.0 \ --port 30000 # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:30000/v1/chat/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "divaspoudel/iol-2026-solver", "messages": [ { "role": "user", "content": "What is the capital of France?" } ] }'Use Docker images
docker run --gpus all \ --shm-size 32g \ -p 30000:30000 \ -v ~/.cache/huggingface:/root/.cache/huggingface \ --env "HF_TOKEN=<secret>" \ --ipc=host \ lmsysorg/sglang:latest \ python3 -m sglang.launch_server \ --model-path "divaspoudel/iol-2026-solver" \ --host 0.0.0.0 \ --port 30000 # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:30000/v1/chat/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "divaspoudel/iol-2026-solver", "messages": [ { "role": "user", "content": "What is the capital of France?" } ] }' - Docker Model Runner
How to use divaspoudel/iol-2026-solver with Docker Model Runner:
docker model run hf.co/divaspoudel/iol-2026-solver
| import os | |
| os.environ["HF_HUB_OFFLINE"] = "1" | |
| os.environ["TRANSFORMERS_OFFLINE"] = "1" | |
| # --- make sure quantization lib is present (PyPI is reachable at run time) --- | |
| import subprocess, sys | |
| def _run_pkg_cmd(args): | |
| for cmd in (["uv", "pip", *args], [sys.executable, "-m", "pip", *args], ["pip", *args]): | |
| try: | |
| if subprocess.run(cmd, capture_output=True, text=True).returncode == 0: | |
| return True | |
| except FileNotFoundError: | |
| continue | |
| return False | |
| try: | |
| import bitsandbytes # noqa: F401 | |
| except ImportError: | |
| _run_pkg_cmd(["install", "-q", "bitsandbytes"]) | |
| # NOTE: torchvision is uninstalled here because it can pull in a conflicting | |
| # pinned torch version when transformers/bitsandbytes resolve dependencies. | |
| # We don't use any vision functionality, so this is safe. | |
| try: | |
| import torchvision # noqa: F401 | |
| _run_pkg_cmd(["uninstall", "-y", "-q", "torchvision"]) | |
| except ImportError: | |
| pass | |
| import re | |
| import json | |
| import time | |
| import math | |
| import os | |
| import sys | |
| import pandas as pd | |
| import torch | |
| from collections import Counter | |
| from transformers import AutoTokenizer, AutoModelForCausalLM, BitsAndBytesConfig | |
| # Safety configuration | |
| DISABLE_RL = os.environ.get("IOL_SAFE_MODE", "").lower() == "true" | |
| if DISABLE_RL: | |
| print("SAFE MODE: RL and HP tuning disabled", flush=True) | |
| MODEL_ID = "." | |
| TIME_LIMIT = 30 * 60 # hard competition limit, seconds | |
| SAFETY_BUFFER = 90 # stop issuing new generations this many seconds before the limit | |
| MAX_NEW_TOKENS = 1800 # increased for scratchpad + rules | |
| NUM_PATHS = 3 # N=3 for test-time scaling | |
| TEMPERATURE = 0.3 # T=0.3 for diverse but coherent paths | |
| TOP_P = 0.9 # default nucleus sampling value (FIX: was only ever set as a | |
| # side-effect global inside apply_tuned_hyperparams and never | |
| # actually read by generate_n_paths) | |
| START = time.time() | |
| def time_left(): | |
| return TIME_LIMIT - (time.time() - START) | |
| print("Loading tokenizer/model...", flush=True) | |
| tok = AutoTokenizer.from_pretrained(MODEL_ID) | |
| bnb_config = BitsAndBytesConfig( | |
| load_in_4bit=True, | |
| bnb_4bit_quant_type="nf4", | |
| bnb_4bit_compute_dtype=torch.float16, | |
| bnb_4bit_use_double_quant=True, | |
| ) | |
| # Load model with OOM protection | |
| try: | |
| model = AutoModelForCausalLM.from_pretrained( | |
| MODEL_ID, | |
| device_map="auto", | |
| torch_dtype=torch.float32, | |
| ).eval() | |
| print("Model loaded successfully", flush=True) | |
| except torch.cuda.OutOfMemoryError as e: | |
| print(f"FATAL: GPU OOM during model load: {e}", flush=True) | |
| # Try CPU fallback | |
| try: | |
| print("Attempting CPU loading...", flush=True) | |
| model = AutoModelForCausalLM.from_pretrained( | |
| MODEL_ID, | |
| device_map="auto", | |
| torch_dtype=torch.float32, | |
| ).eval() | |
| except Exception as e2: | |
| print(f"FATAL: Could not load model even on CPU: {e2}", flush=True) | |
| # Create error submission | |
| with open("submission.csv", "w") as f: | |
| f.write("id,pred,explanation\n") | |
| sys.exit(1) | |
| print(f"Model loaded in {time.time() - START:.1f}s", flush=True) | |
| # Robust data loading with error handling | |
| try: | |
| if not os.path.exists("/tmp/data/test.csv"): | |
| print("ERROR: test.csv not found at /tmp/data/test.csv", flush=True) | |
| # Create empty submission as fallback | |
| with open("submission.csv", "w") as f: | |
| f.write("id,pred,explanation\n") | |
| sys.exit(1) | |
| df = pd.read_csv("/tmp/data/test.csv", dtype=str).fillna("") | |
| if len(df) == 0: | |
| print("WARNING: test.csv is empty", flush=True) | |
| # Still create valid submission | |
| pd.DataFrame(columns=["id", "pred", "explanation"]).to_csv("submission.csv", index=False) | |
| sys.exit(0) | |
| print(f"Loaded {len(df)} test items", flush=True) | |
| except Exception as e: | |
| print(f"FATAL: Could not load test.csv: {e}", flush=True) | |
| # Create error submission | |
| with open("submission.csv", "w") as f: | |
| f.write("id,pred,explanation\n") | |
| sys.exit(1) | |
| ITEM_RE = re.compile(r"(?m)^\s*(\d+)\.\s") | |
| # ============================================================================= | |
| # 1. DYNAMIC TASK ROUTER | |
| # Analyzes: task_type, eval_type, item count (K) | |
| # ============================================================================= | |
| TASK_PROFILES = { | |
| "translation": { | |
| "eval_metric": "chrF", | |
| "format_hint": "Translate to the target language preserving diacritics and orthography.", | |
| "grammar_focus": "morphophonology, tense, case, agreement", | |
| }, | |
| "match_letters": { | |
| "eval_metric": "exact", | |
| "format_hint": "Answer with only the option letter (e.g. A, B, C) for each item.", | |
| "grammar_focus": "phonological rules, allophony, orthographic patterns", | |
| }, | |
| "fill_blanks": { | |
| "eval_metric": "chrF", | |
| "format_hint": "Fill the blank with the correct inflected/corrected form.", | |
| "grammar_focus": "inflection, derivation, agreement, sandhi", | |
| }, | |
| "text_to_num": { | |
| "eval_metric": "exact", | |
| "format_hint": "Answer with the value written in digits (e.g. 285).", | |
| "grammar_focus": "numeral systems, base systems, place value", | |
| }, | |
| "num_to_text": { | |
| "eval_metric": "chrF", | |
| "format_hint": "Answer with the number written out in words in the task language.", | |
| "grammar_focus": "numeral morphology, ordinals, cardinals", | |
| }, | |
| } | |
| def count_items(query: str) -> int: | |
| """Extract K = number of items in the query. | |
| FIX: previously used len(set(nums)) which under-counts whenever a numeric | |
| label repeats or numbering is non-contiguous (e.g. duplicated "1." lines, | |
| grouped sub-items). Using the max item index found is robust to that, | |
| since IOL-style queries number items sequentially starting at 1. | |
| """ | |
| nums = ITEM_RE.findall(query) | |
| if nums: | |
| return max(int(n) for n in nums) | |
| # fallback: count non-empty lines | |
| return max(1, len([l for l in query.splitlines() if l.strip()])) | |
| def get_task_profile(task_type: str) -> dict: | |
| """Route to task-specific configuration.""" | |
| return TASK_PROFILES.get(task_type, TASK_PROFILES["translation"]) | |
| # ============================================================================= | |
| # 2. FEW-SHOT EXAMPLES: Filled-out Grammar Scratchpads | |
| # Shows the model HOW to reason, not just WHAT to do | |
| # ============================================================================= | |
| FEW_SHOT_EXAMPLES = """ | |
| === EXAMPLE 1: Simple Translation === | |
| DATA: | |
| 1. mo ka | I go | |
| 2. ti ka | you go | |
| 3. mo tuʔ | I sleep | |
| 4. ti tuʔ | you sleep | |
| QUERY: Translate into English: | |
| 1. ti tuʔ | |
| 2. mo ka | |
| 【GRAMMAR SCRATCHPAD】 | |
| 1. DATA INVENTORY: | |
| - mo: "I" (appears in sentences 1, 3) | |
| - ti: "you" (appears in sentences 2, 4) | |
| - ka: "go" (appears in sentences 1, 2) | |
| - tuʔ: "sleep" (appears in sentences 3, 4) | |
| 2. MORPHEME SEGMENTATION: | |
| - All words are monomorphemic (no affixes detected) | |
| - Structure: [PRONOUN] + [VERB] | |
| 3. PARADIGM MAPPING: | |
| - Pronouns: mo = 1st person singular, ti = 2nd person singular | |
| - Verbs: ka = "go", tuʔ = "sleep" (no inflection seen) | |
| 4. CONSTRAINT IDENTIFICATION: | |
| - No agreement marking on verbs | |
| - Word order: PRONOUN + VERB | |
| 5. VERIFICATION: | |
| - Sentence 1: mo ka = "I go" - matches data | |
| - Sentence 2: ti ka = "you go" - matches data | |
| - Pattern is consistent | |
| 6. APPLICATION: | |
| - ti tuʔ = "you + sleep" = "you sleep" | |
| - mo ka = "I + go" = "I go" | |
| RULES: | |
| - RULE 1: mo -> "I" / _ (1st person singular pronoun) | |
| - RULE 2: ti -> "you" / _ (2nd person singular pronoun) | |
| - RULE 3: ka -> "go" / _ (uninflected verb) | |
| - RULE 4: tuʔ -> "sleep" / _ (uninflected verb) | |
| - RULE 5: [PRONOUN] + [VERB] -> [PRONOUN] + [VERB] | |
| EXPLANATION: The language uses independent pronouns and invariant verbs. Word order is PRONOUN-VERB. | |
| FINAL: ["you sleep", "I go"] | |
| === EXAMPLE 2: Morphological Rule Detection === | |
| DATA: | |
| 1. book | one book | |
| 2. books | two books | |
| 3. box | one box | |
| 4. boxes | two boxes | |
| QUERY: Complete the translation: | |
| 1. cat | one cat | |
| 2. cats | ___ | |
| 【GRAMMAR SCRATCHPAD】 | |
| 1. DATA INVENTORY: | |
| - "book" appears as "book" (one) and "books" (two) | |
| - "box" appears as "box" (one) and "boxes" (two) | |
| - "one" marks singular, "two" marks plural | |
| 2. MORPHEME SEGMENTATION: | |
| - books = book + -s | |
| - boxes = box + -es | |
| - Singular forms lack suffix | |
| 3. PARADIGM MAPPING: | |
| - book ~ books: add -s | |
| - box ~ boxes: add -es | |
| - Pattern: plural = stem + suffix | |
| 4. CONSTRAINT IDENTIFICATION: | |
| - ALTERNATION: -s vs -es depends on stem-final sound | |
| - "book" ends in /k/ (non-sibilant) -> -s | |
| - "box" ends in /ks/ (sibilant) -> -es | |
| - RULE: Use -es after sibilants (/s/, /z/, /ʃ/, /tʃ/, /ks/, etc.) | |
| 5. VERIFICATION: | |
| - "book" + "s" = "books" - matches data | |
| - "box" + "es" = "boxes" - matches data | |
| - Rule holds | |
| 6. APPLICATION: | |
| - "cat" ends in /t/ (non-sibilant) -> plural = "cats" | |
| - "two" + "cats" = "two cats" | |
| RULES: | |
| - RULE 1: NOUN_sg -> NOUN_pl / [two] (context: plural number) | |
| - RULE 2: X -> X-s / _# (default plural: add -s) | |
| - RULE 3: X[sibilant] -> X-es / _# (sibilant plural: add -es) | |
| - RULE 4: Sibilant set: {s, z, ʃ, ʒ, tʃ, dʒ, ks, gz} | |
| - RULE 5: [two] + [NOUN_pl] -> "two" + [NOUN_pl] | |
| EXPLANATION: Plural formation uses -s by default. After sibilant sounds (s, z, sh, ch, x, etc.), use -es instead. "cat" ends in /t/, so regular -s plural applies. | |
| FINAL: ["one cat", "two cats"] | |
| """ | |
| GRAMMAR_SCRATCHPAD_TEMPLATE = """Before answering, you MUST work through this deduction process: | |
| 【GRAMMAR SCRATCHPAD】 | |
| 1. DATA INVENTORY: List forms and glosses explicitly shown | |
| 2. MORPHEME SEGMENTATION: Break forms into morphemes, mark boundaries with - | |
| 3. PARADIGM MAPPING: Sketch inflectional categories (person, number, tense, case) | |
| 4. CONSTRAINT IDENTIFICATION: Phonological rules (sandhi, harmony, mutation), morphological rules (affixation, stem change) | |
| 5. VERIFICATION: Test hypothesis against exceptions | |
| 6. APPLICATION: Apply rules to query items step by step | |
| Then provide: | |
| RULES: List explicit rules in format "RULE N: [input] -> [output] / [condition]" | |
| EXPLANATION: Summary of grammatical analysis | |
| FINAL: JSON array with answers | |
| """ | |
| BASE_SYSTEM_PROMPT = ( | |
| "You are an expert linguist competing in the International Linguistics Olympiad. " | |
| "Analyze linguistic data, deduce grammar rules, and solve problems.\n\n" | |
| + FEW_SHOT_EXAMPLES + "\n\n" | |
| + GRAMMAR_SCRATCHPAD_TEMPLATE + "\n\n" | |
| "LINGUISTIC PHENOMENA TO CONSIDER:\n" | |
| "- Phonology: assimilation, dissimilation, deletion, epenthesis, metathesis, lenition, fortition\n" | |
| "- Morphology: prefixation, suffixation, infixation, circumfixation, reduplication, ablaut, suppletion\n" | |
| "- Phonotactics: consonant clusters, vowel constraints, syllable structure\n" | |
| "- Sandhi: external (word-boundary), internal (morpheme-boundary)\n" | |
| "- Syntax: word order (SOV, SVO, VSO), agreement patterns, case marking\n\n" | |
| "Respond with these sections:\n" | |
| "1. 【GRAMMAR SCRATCHPAD】 (complete all 6 steps with specific analysis)\n" | |
| "2. RULES: Explicit rewrite rules like:\n" | |
| " - RULE 1: [morpheme A] + [morpheme B] -> [result] / [environment]\n" | |
| " - RULE 2: X -> Y / _(condition) (phonological rule)\n" | |
| " - RULE 3: [CATEGORY] -> [translation/meaning]\n" | |
| "3. EXPLANATION: Clear summary of reasoning and rules\n" | |
| "4. FINAL: JSON array [\"answer1\", \"answer2\", ...] with exactly one string per query item" | |
| ) | |
| def build_router_aware_prompt(context: str, query: str, task_type: str, eval_type: str, k: int) -> str: | |
| """Build a prompt enriched with task routing information.""" | |
| profile = get_task_profile(task_type) | |
| header = f"""【TASK PROFILE】 | |
| Task Type: {task_type} | |
| Evaluation: {eval_type} (using {profile['eval_metric']}) | |
| Number of Items: {k} | |
| Focus: {profile['grammar_focus']} | |
| """ | |
| hint = profile["format_hint"] | |
| return f"{header}\n---\nDATA:\n{context.strip()}\n\nQUERY ({k} items):\n{query.strip()}\n\nFormat: {hint}" | |
| # ============================================================================= | |
| # 3. TEST-TIME SCALING (Self-Consistency Engine) | |
| # ============================================================================= | |
| def generate_n_paths(prompt: str, n: int = None, temperature: float = None, top_p: float = None) -> list: | |
| """Generate N diverse reasoning paths using sampling. | |
| FIX: now reads the (possibly tuned) global NUM_PATHS/TEMPERATURE/TOP_P/ | |
| MAX_NEW_TOKENS at call time via defaults, and top_p is actually threaded | |
| through to model.generate() instead of being hardcoded to 0.9. | |
| """ | |
| if n is None: | |
| n = NUM_PATHS | |
| if temperature is None: | |
| temperature = TEMPERATURE | |
| if top_p is None: | |
| top_p = TOP_P | |
| messages = [ | |
| {"role": "system", "content": BASE_SYSTEM_PROMPT}, | |
| {"role": "user", "content": prompt}, | |
| ] | |
| model_inputs = tok.apply_chat_template( | |
| messages, add_generation_prompt=True, return_tensors="pt" | |
| ) | |
| ids = model_inputs['input_ids'].to(model.device) | |
| paths = [] | |
| for path_idx in range(n): | |
| try: | |
| with torch.no_grad(): | |
| out = model.generate( | |
| ids, | |
| max_new_tokens=MAX_NEW_TOKENS, | |
| do_sample=True, | |
| temperature=temperature, | |
| top_p=top_p, | |
| pad_token_id=tok.eos_token_id, | |
| ) | |
| text = tok.decode(out[0][ids.shape[-1]:], skip_special_tokens=True).strip() | |
| paths.append(text) | |
| print(f" Path {path_idx+1}/{n} generated", flush=True) | |
| except Exception as e: | |
| print(f" Path {path_idx+1} failed: {e}", flush=True) | |
| paths.append("") | |
| return paths | |
| def chrF_score(candidate: str, reference: str) -> float: | |
| """Compute character n-gram F-score (simplified chrF).""" | |
| if not candidate or not reference: | |
| return 0.0 | |
| def get_ngrams(s, n): | |
| s = s.lower() | |
| return [s[i:i+n] for i in range(len(s) - n + 1)] | |
| total_f = 0.0 | |
| for n in range(2, 7): | |
| cand_ngrams = Counter(get_ngrams(candidate, n)) | |
| ref_ngrams = Counter(get_ngrams(reference, n)) | |
| overlapping = sum((cand_ngrams & ref_ngrams).values()) | |
| precision = overlapping / max(sum(cand_ngrams.values()), 1) | |
| recall = overlapping / max(sum(ref_ngrams.values()), 1) | |
| if precision + recall > 0: | |
| f = 2 * precision * recall / (precision + recall) | |
| total_f += f | |
| return total_f / 6.0 | |
| def consensus_vote(path_answers: list, eval_metric: str) -> tuple: | |
| """ | |
| Aggregate N paths using consensus voting. | |
| Returns (consensus_answer, confidence, explanation) | |
| """ | |
| if not path_answers or len(path_answers) == 0: | |
| return [], 0.0, "No paths generated" | |
| path_answers = [ans for ans in path_answers if ans] | |
| if not path_answers: | |
| return [], 0.0, "All paths failed" | |
| k = len(path_answers[0]) | |
| consensus = [] | |
| confidences = [] | |
| for item_idx in range(k): | |
| item_answers = [] | |
| for path in path_answers: | |
| if item_idx < len(path): | |
| item_answers.append(path[item_idx]) | |
| if not item_answers: | |
| consensus.append("") | |
| confidences.append(0.0) | |
| continue | |
| if eval_metric == "exact": | |
| answer_counts = Counter(item_answers) | |
| best_answer, count = answer_counts.most_common(1)[0] | |
| confidence = count / len(item_answers) | |
| consensus.append(best_answer) | |
| confidences.append(confidence) | |
| else: | |
| best_answer = item_answers[0] | |
| best_score = 0.0 | |
| for candidate in item_answers: | |
| score = sum(chrF_score(candidate, other) for other in item_answers) / len(item_answers) | |
| if score > best_score: | |
| best_score = score | |
| best_answer = candidate | |
| agreeing = sum(1 for ans in item_answers if chrF_score(best_answer, ans) > 0.8) | |
| confidence = agreeing / len(item_answers) | |
| consensus.append(best_answer) | |
| confidences.append(confidence) | |
| avg_confidence = sum(confidences) / len(confidences) if confidences else 0.0 | |
| explanation_parts = [ | |
| f"Generated {len(path_answers)} paths with T={TEMPERATURE}", | |
| f"Consensus ({eval_metric}) confidence: {avg_confidence:.0%}", | |
| ] | |
| return consensus, avg_confidence, " | ".join(explanation_parts) | |
| def extract_answers_from_text(text: str, expected_k: int) -> list: | |
| """Extract the JSON answer array from model output. | |
| FIX: the generic bracket-scanning fallback now scans matches in reverse | |
| (closest to the end of the text first), since FINAL: is always last in | |
| the expected output format, and earlier [...] occurrences (e.g. sibilant | |
| sets, category lists in RULES) were sometimes matched first and returned | |
| the wrong array when the FINAL: regex failed to match (e.g. truncated | |
| generation cut off before the closing bracket). | |
| """ | |
| # Try FINAL: block first | |
| m = re.search(r"FINAL:\s*(\[.*?\])", text, re.DOTALL | re.IGNORECASE) | |
| if m: | |
| try: | |
| parsed = json.loads(m.group(1)) | |
| if isinstance(parsed, list): | |
| return [str(x) for x in parsed] | |
| except Exception: | |
| pass | |
| # Look for any JSON array, preferring the one closest to the end of the | |
| # text (FINAL: is always the last section in the expected format). | |
| matches = list(re.finditer(r"\[.*?\]", text, re.DOTALL)) | |
| for m2 in reversed(matches): | |
| try: | |
| parsed = json.loads(m2.group(0)) | |
| if isinstance(parsed, list) and parsed: | |
| return [str(x) for x in parsed] | |
| except Exception: | |
| continue | |
| # Fallback: extract lines after FINAL: | |
| tail = text.split("FINAL:")[-1] | |
| lines = [re.sub(r"^\s*\d+[\.\)\]]\s*", "", l).strip(' \t\"\'') | |
| for l in tail.splitlines() if l.strip()] | |
| lines = [l for l in lines if l] | |
| return lines | |
| def extract_rules(text: str) -> list: | |
| """Extract explicit rules from RULES: section.""" | |
| rules = [] | |
| # Look for RULES: section | |
| rules_match = re.search(r"RULES:(.+?)(?=EXPLANATION:|FINAL:|$)", text, re.DOTALL | re.IGNORECASE) | |
| if rules_match: | |
| rules_text = rules_match.group(1) | |
| # Extract individual rule lines starting with - or RULE | |
| for line in rules_text.split('\n'): | |
| line = line.strip() | |
| if line and (line.startswith('-') or line.startswith('RULE') or | |
| (len(line) > 10 and '->' in line)): | |
| # Clean up the rule | |
| rule = re.sub(r'^[-•*]\s*', '', line).strip() | |
| if rule and len(rule) > 5: | |
| rules.append(rule) | |
| return rules | |
| def extract_explanation(text: str) -> str: | |
| """Extract the explanation portion from output.""" | |
| m = re.search(r"EXPLANATION:\s*(.+?)(?:RULES:|FINAL:|$)", text, re.DOTALL | re.IGNORECASE) | |
| if m: | |
| return m.group(1).strip()[:800] | |
| return "" | |
| # ============================================================================= | |
| # HYPERPARAMETER TUNING ENGINE | |
| # Optimizes: NUM_PATHS, TEMPERATURE, MAX_NEW_TOKENS, TOP_P | |
| # Uses small validation set to find best configuration | |
| # ============================================================================= | |
| HYPERPARAM_SPACE = { | |
| "num_paths": [2, 3, 4, 5], # N: number of reasoning paths | |
| "temperature": [0.1, 0.2, 0.3, 0.4, 0.5], # T: sampling diversity | |
| "top_p": [0.85, 0.9, 0.95, 0.99], # Nucleus sampling | |
| "max_tokens_factor": [1.0, 1.2, 1.5], # Multiplier for base tokens | |
| } | |
| # Validation set for hyperparameter tuning (subset of typical IOL problems) | |
| VALIDATION_EXAMPLES = [ | |
| { | |
| "id": "val-001", | |
| "task_type": "translation", | |
| "context": "1. áaka | I see\n2. tíika | you see\n3. áatʃi | I walk\n4. tíitʃi | you walk", | |
| "query": "Translate:\n1. áatʃi\n2. tíika", | |
| "expected": ["I walk", "you see"], | |
| "eval_metric": "chrF", | |
| }, | |
| { | |
| "id": "val-002", | |
| "task_type": "match_letters", | |
| "context": "1. atu A. water\n2. keno B. fire\n3. suna C. sun\n4. mizu D. east\n5. umi E. sea", | |
| "query": "Match:\n1. atu\n2. keno\n3. suna", | |
| "expected": ["A", "B", "C"], | |
| "eval_metric": "exact", | |
| }, | |
| ] | |
| def evaluate_hyperparams(config: dict, val_examples: list, max_evals: int = 2) -> dict: | |
| """ | |
| Evaluate a hyperparameter configuration on validation examples. | |
| Returns: {"score": float, "avg_time": float, "consistency": float} | |
| """ | |
| n_paths = config["num_paths"] | |
| temp = config["temperature"] | |
| top_p = config["top_p"] | |
| max_tokens = int(MAX_NEW_TOKENS * config["max_tokens_factor"]) | |
| total_score = 0.0 | |
| total_time = 0.0 | |
| consistency_scores = [] | |
| # Only evaluate on first max_evals examples for speed | |
| for example in val_examples[:max_evals]: | |
| start_t = time.time() | |
| # Generate N paths | |
| messages = [ | |
| {"role": "system", "content": BASE_SYSTEM_PROMPT}, | |
| {"role": "user", "content": example["context"] + "\n\n" + example["query"]}, | |
| ] | |
| model_inputs = tok.apply_chat_template( | |
| messages, add_generation_prompt=True, return_tensors="pt" | |
| ) | |
| ids = model_inputs['input_ids'].to(model.device) | |
| paths = [] | |
| for _ in range(n_paths): | |
| try: | |
| with torch.no_grad(): | |
| out = model.generate( | |
| ids, | |
| max_new_tokens=max_tokens, | |
| do_sample=True, | |
| temperature=temp, | |
| top_p=top_p, | |
| pad_token_id=tok.eos_token_id, | |
| ) | |
| text = tok.decode(out[0][ids.shape[-1]:], skip_special_tokens=True).strip() | |
| paths.append(text) | |
| except Exception: | |
| paths.append("") | |
| gen_time = time.time() - start_t | |
| total_time += gen_time | |
| # Extract answers | |
| path_answers = [] | |
| for path in paths: | |
| if path: | |
| answers = extract_answers_from_text(path, len(example["expected"])) | |
| answers = enforce_k_length(answers, len(example["expected"])) | |
| path_answers.append(answers) | |
| if not path_answers: | |
| continue | |
| # Calculate consensus and score | |
| if example["eval_metric"] == "exact": | |
| # For exact match: check if consensus matches expected | |
| answers_list = [tuple(a) for a in path_answers if a] | |
| if answers_list: | |
| most_common = Counter(answers_list).most_common(1)[0][0] | |
| consensus = list(most_common) | |
| score = sum(1 for i, exp in enumerate(example["expected"]) | |
| if i < len(consensus) and consensus[i] == exp) / len(example["expected"]) | |
| # Consistency = agreement among paths | |
| consistency = Counter(answers_list).most_common(1)[0][1] / len(answers_list) | |
| else: | |
| score = 0.0 | |
| consistency = 0.0 | |
| else: | |
| # For chrF: use fuzzy matching | |
| consensus, _, _ = consensus_vote(path_answers, "chrF") | |
| score = sum(chrF_score(consensus[i], example["expected"][i]) | |
| for i in range(min(len(consensus), len(example["expected"])))) / len(example["expected"]) | |
| # Consistency based on path agreement | |
| consistency = sum( | |
| sum(chrF_score(a1, a2) for a2 in path_answers) / len(path_answers) | |
| for a1 in path_answers | |
| ) / len(path_answers) if path_answers else 0.0 | |
| total_score += score | |
| consistency_scores.append(consistency) | |
| n_evaluated = min(max_evals, len(val_examples)) | |
| avg_score = total_score / n_evaluated if n_evaluated > 0 else 0.0 | |
| avg_time = total_time / n_evaluated if n_evaluated > 0 else 0.0 | |
| avg_consistency = sum(consistency_scores) / len(consistency_scores) if consistency_scores else 0.0 | |
| return { | |
| "score": avg_score, | |
| "avg_time": avg_time, | |
| "consistency": avg_consistency, | |
| } | |
| def grid_search_hyperparams(val_examples: list, max_configs: int = 8) -> dict: | |
| """ | |
| Grid search over hyperparameter space. | |
| Tests promising configurations and returns the best one. | |
| """ | |
| print("\n=== HYPERPARAMETER TUNING ===", flush=True) | |
| print(f"Testing configurations on {len(val_examples)} validation examples...", flush=True) | |
| # Priority order based on typical IOL performance | |
| test_configs = [ | |
| {"num_paths": 3, "temperature": 0.2, "top_p": 0.9, "max_tokens_factor": 1.0}, | |
| {"num_paths": 3, "temperature": 0.3, "top_p": 0.9, "max_tokens_factor": 1.0}, | |
| {"num_paths": 4, "temperature": 0.2, "top_p": 0.95, "max_tokens_factor": 1.2}, | |
| {"num_paths": 2, "temperature": 0.1, "top_p": 0.85, "max_tokens_factor": 1.0}, | |
| {"num_paths": 5, "temperature": 0.3, "top_p": 0.95, "max_tokens_factor": 1.2}, | |
| {"num_paths": 3, "temperature": 0.4, "top_p": 0.9, "max_tokens_factor": 1.0}, | |
| {"num_paths": 4, "temperature": 0.2, "top_p": 0.9, "max_tokens_factor": 1.5}, | |
| {"num_paths": 3, "temperature": 0.2, "top_p": 0.99, "max_tokens_factor": 1.2}, | |
| ][:max_configs] | |
| results = [] | |
| for i, config in enumerate(test_configs): | |
| # FIX: also bail out of tuning itself if time is getting short, so | |
| # HP search can't eat into the main processing budget unbounded. | |
| if time_left() < TIME_LIMIT * 0.6: | |
| print(f" Stopping HP search early: time_left={time_left():.0f}s", flush=True) | |
| break | |
| print(f"\nConfig {i+1}/{len(test_configs)}: N={config['num_paths']}, T={config['temperature']}, top_p={config['top_p']}", flush=True) | |
| metrics = evaluate_hyperparams(config, val_examples) | |
| results.append((config, metrics)) | |
| print(f" Score: {metrics['score']:.3f}, Consistency: {metrics['consistency']:.3f}, Time: {metrics['avg_time']:.2f}s", flush=True) | |
| if not results: | |
| # Nothing evaluated (ran out of time immediately) - fall back to defaults | |
| return test_configs[0] | |
| # Select best configuration based on combined score | |
| # Weight: accuracy 60%, consistency 30%, speed 10% | |
| def combined_score(config, metrics): | |
| # Normalize time (lower is better, assume max 30s per example) | |
| time_score = max(0, 1 - metrics['avg_time'] / 30.0) | |
| return (0.6 * metrics['score'] + | |
| 0.3 * metrics['consistency'] + | |
| 0.1 * time_score) | |
| best_config, best_metrics = max(results, key=lambda x: combined_score(x[0], x[1])) | |
| print(f"\n=== BEST CONFIGURATION ===", flush=True) | |
| print(f"NUM_PATHS: {best_config['num_paths']}", flush=True) | |
| print(f"TEMPERATURE: {best_config['temperature']}", flush=True) | |
| print(f"TOP_P: {best_config['top_p']}", flush=True) | |
| print(f"MAX_TOKENS_FACTOR: {best_config['max_tokens_factor']}", flush=True) | |
| print(f"Expected Score: {best_metrics['score']:.3f}", flush=True) | |
| return best_config | |
| # Global hyperparameters (will be set by tuning or use defaults) | |
| TUNED_NUM_PATHS = None | |
| TUNED_TEMPERATURE = None | |
| TUNED_TOP_P = None | |
| TUNED_MAX_TOKENS_FACTOR = None | |
| def apply_tuned_hyperparams(config: dict = None): | |
| """Apply tuned hyperparameters to global variables. | |
| FIX: previously TOP_P and the scaled MAX_NEW_TOKENS were computed and | |
| printed but never actually written back to the globals that | |
| generate_n_paths() reads, so all tuning was a no-op. Both are now | |
| applied. Also guards against re-scaling MAX_NEW_TOKENS more than once | |
| if this function is ever called twice. | |
| """ | |
| global NUM_PATHS, TEMPERATURE, TOP_P, MAX_NEW_TOKENS | |
| global TUNED_NUM_PATHS, TUNED_TEMPERATURE, TUNED_TOP_P, TUNED_MAX_TOKENS_FACTOR | |
| global _BASE_MAX_NEW_TOKENS | |
| if '_BASE_MAX_NEW_TOKENS' not in globals(): | |
| _BASE_MAX_NEW_TOKENS = MAX_NEW_TOKENS | |
| if config is None: | |
| # Use default/baseline configuration | |
| config = { | |
| "num_paths": NUM_PATHS, | |
| "temperature": TEMPERATURE, | |
| "top_p": TOP_P, | |
| "max_tokens_factor": 1.0, | |
| } | |
| TUNED_NUM_PATHS = config["num_paths"] | |
| TUNED_TEMPERATURE = config["temperature"] | |
| TUNED_TOP_P = config.get("top_p", TOP_P) | |
| TUNED_MAX_TOKENS_FACTOR = config.get("max_tokens_factor", 1.0) | |
| # Update the globals actually used by generate_n_paths() and friends | |
| NUM_PATHS = TUNED_NUM_PATHS | |
| TEMPERATURE = TUNED_TEMPERATURE | |
| TOP_P = TUNED_TOP_P | |
| MAX_NEW_TOKENS = int(_BASE_MAX_NEW_TOKENS * TUNED_MAX_TOKENS_FACTOR) | |
| print(f"\nApplied tuned hyperparameters:", flush=True) | |
| print(f" NUM_PATHS: {NUM_PATHS}", flush=True) | |
| print(f" TEMPERATURE: {TEMPERATURE}", flush=True) | |
| print(f" TOP_P: {TOP_P}", flush=True) | |
| print(f" MAX_NEW_TOKENS: {MAX_NEW_TOKENS}", flush=True) | |
| # ============================================================================= | |
| # REINFORCEMENT LEARNING COMPONENT | |
| # Online learning from test samples using self-consistency as reward | |
| # ============================================================================= | |
| class ReinforcementLearner: | |
| """ | |
| Online RL that samples from test data and reinforces successful patterns. | |
| Uses self-consistency and confidence as reward signals. | |
| """ | |
| def __init__(self): | |
| self.learned_patterns = {} # pattern -> success_score | |
| self.morpheme_inventory = {} # morpheme -> {translations, contexts} | |
| self.rule_confidence = {} # rule -> confidence_score | |
| self.paradigm_library = {} # pattern_type -> paradigms | |
| self.iteration_rewards = [] # Track reward over iterations | |
| def sample_test_data(self, df: pd.DataFrame, n_samples: int = 3) -> pd.DataFrame: | |
| """ | |
| Sample representative items from test set for RL training. | |
| Strategy: diverse task types, one (or more) sample per task type, | |
| trimmed to n_samples total. | |
| FIX: previous implementation concatenated per-task-type samples then | |
| applied .head(n_samples), which meant only the first task type(s) in | |
| groupby-iteration order ever survived the final truncation - i.e. it | |
| did NOT actually guarantee cross-task-type diversity as intended. | |
| This version samples a balanced number per group first, then trims | |
| with a random sample instead of a positional head(). | |
| """ | |
| if len(df) <= n_samples: | |
| return df.copy() | |
| groups = list(df.groupby('task_type')) | |
| if not groups: | |
| return df.sample(n_samples, random_state=0) | |
| per_group = max(1, n_samples // len(groups)) | |
| picked = [] | |
| for _, type_df in groups: | |
| type_df = type_df.copy() | |
| type_df['context_len'] = type_df['context'].str.len() | |
| median_len = type_df['context_len'].median() | |
| type_df['complexity_score'] = (type_df['context_len'] - median_len).abs() | |
| type_df = type_df.sort_values('complexity_score') | |
| picked.append(type_df.head(per_group)) | |
| sampled = pd.concat(picked) | |
| if len(sampled) > n_samples: | |
| sampled = sampled.sample(n_samples, random_state=0) | |
| return sampled.drop(columns=['context_len', 'complexity_score'], errors='ignore') | |
| def compute_reward(self, paths: list, consensus: list, confidence: float, | |
| expected_k: int, generation_time: float) -> dict: | |
| """ | |
| Compute multi-faceted reward for RL training. | |
| Reward components: | |
| - consistency_reward: Agreement between paths (self-consistency) | |
| - confidence_reward: Model confidence in its answer | |
| - format_reward: Proper JSON formatting and K-length | |
| - efficiency_reward: Reasonable generation time | |
| """ | |
| # Extract answers from all paths | |
| path_answers = [] | |
| for path in paths: | |
| if path: | |
| answers = extract_answers_from_text(path, expected_k) | |
| if answers: | |
| path_answers.append(tuple(answers)) | |
| # Consistency reward: pairwise agreement | |
| if len(path_answers) >= 2: | |
| agreements = 0 | |
| total_pairs = 0 | |
| for i in range(len(path_answers)): | |
| for j in range(i+1, len(path_answers)): | |
| # Compare answer similarity | |
| if len(path_answers[i]) == len(path_answers[j]): | |
| matches = sum(1 for a, b in zip(path_answers[i], path_answers[j]) | |
| if chrF_score(a, b) > 0.8) | |
| agreements += matches / len(path_answers[i]) | |
| total_pairs += 1 | |
| consistency_reward = agreements / total_pairs if total_pairs > 0 else 0.0 | |
| else: | |
| consistency_reward = 0.0 | |
| # Confidence reward (already computed) | |
| confidence_reward = confidence | |
| # Format reward: Did we get valid K-length outputs? | |
| format_reward = 1.0 if len(consensus) == expected_k and all(consensus) else 0.5 | |
| # Efficiency reward: Prefer faster generations (normalize to 30s target) | |
| efficiency_reward = max(0, 1 - generation_time / 30.0) | |
| # Combined reward (weighted) | |
| total_reward = ( | |
| 0.35 * consistency_reward + | |
| 0.35 * confidence_reward + | |
| 0.20 * format_reward + | |
| 0.10 * efficiency_reward | |
| ) | |
| return { | |
| "total": total_reward, | |
| "consistency": consistency_reward, | |
| "confidence": confidence_reward, | |
| "format": format_reward, | |
| "efficiency": efficiency_reward, | |
| } | |
| def extract_patterns(self, paths: list, consensus: list, task_type: str) -> list: | |
| """ | |
| Extract successful/reusable patterns from generated paths. | |
| """ | |
| patterns = [] | |
| for path in paths: | |
| if not path: | |
| continue | |
| # Extract morphemes from scratchpad | |
| morpheme_section = re.search( | |
| r'2\.\s*MORPHEME SEGMENTATION:(.+?)(?=3\.|PARADIGM|RULES|EXPLANATION|$)', | |
| path, re.DOTALL | re.IGNORECASE | |
| ) | |
| if morpheme_section: | |
| # Look for morpheme mappings like "word = morpheme1 + morpheme2" | |
| mappings = re.findall( | |
| r'([\w\-ʔʼʰʱʲʷː̥̩̪̯̃͡ɑ-ʯḀ-ẞ]+)\s*=\s*([\w\-ʔʼʰʱʲʷː̥̩̪̯̃͡ɑ-ʯḀ-ẞ]+)\s*\+\s*([\w\-ʔʼʰʱʲʷː̥̩̪̯̃͡ɑ-ʯḀ-ẞ]+)', | |
| morpheme_section.group(1) | |
| ) | |
| for mapping in mappings: | |
| pattern_key = f"SEG:{mapping[0]}->{mapping[1]}+{mapping[2]}" | |
| patterns.append(("morpheme_segmentation", pattern_key, mapping)) | |
| # Extract rules from RULES section | |
| rules = extract_rules(path) | |
| for rule in rules: | |
| rule_type = self._classify_rule(rule, task_type) | |
| patterns.append(("rule", rule_type, rule)) | |
| # Extract paradigm patterns | |
| paradigm_section = re.search( | |
| r'3\.\s*PARADIGM MAPPING:(.+?)(?=4\.|CONSTRAINT|RULES|EXPLANATION|$)', | |
| path, re.DOTALL | re.IGNORECASE | |
| ) | |
| if paradigm_section: | |
| # Look for category mappings | |
| categories = re.findall( | |
| r'(1st|2nd|3rd|singular|plural|present|past|nominative|accusative)\s*[=:]\s*([\w\-]+)', | |
| paradigm_section.group(1), re.IGNORECASE | |
| ) | |
| if categories: | |
| patterns.append(("paradigm", task_type, categories)) | |
| return patterns | |
| def _classify_rule(self, rule: str, task_type: str) -> str: | |
| """Classify rule type for organizing the library.""" | |
| rule_lower = rule.lower() | |
| if any(x in rule_lower for x in ['phonol', 'consonant', 'vowel', 'sound', 'harmony', 'sandhi']): | |
| return f"{task_type}:phonological" | |
| elif any(x in rule_lower for x in ['prefix', 'suffix', 'infix', 'affix', 'morpheme']): | |
| return f"{task_type}:morphological" | |
| elif any(x in rule_lower for x in ['word order', 'synta', 'sov', 'svo', 'agreement']): | |
| return f"{task_type}:syntactic" | |
| else: | |
| return f"{task_type}:general" | |
| def update_learned_patterns(self, patterns: list, reward: float): | |
| """ | |
| Update pattern library with exponential moving average. | |
| """ | |
| alpha = 0.3 # Learning rate | |
| for pattern_type, pattern_key, pattern_value in patterns: | |
| key = f"{pattern_type}:{pattern_key}" | |
| if key not in self.learned_patterns: | |
| self.learned_patterns[key] = { | |
| "type": pattern_type, | |
| "key": pattern_key, | |
| "value": pattern_value, | |
| "success_score": reward, | |
| "count": 1, | |
| } | |
| else: | |
| # Exponential moving average update | |
| old_score = self.learned_patterns[key]["success_score"] | |
| self.learned_patterns[key]["success_score"] = ( | |
| (1 - alpha) * old_score + alpha * reward | |
| ) | |
| self.learned_patterns[key]["count"] += 1 | |
| def build_rl_enhanced_prompt(self, base_prompt: str, task_type: str) -> str: | |
| """ | |
| Enhance prompt with learned patterns relevant to this task type. | |
| """ | |
| enhancements = [] | |
| # Add high-confidence patterns for this task type | |
| relevant_patterns = [ | |
| p for k, p in self.learned_patterns.items() | |
| if task_type in str(p.get("key", "")) and p["success_score"] > 0.5 | |
| ] | |
| # Sort by success score | |
| relevant_patterns.sort(key=lambda x: x["success_score"], reverse=True) | |
| if relevant_patterns: | |
| enhancements.append("\n【LEARNED PATTERNS FROM TEST DATA】") | |
| enhancements.append("Based on analysis of similar problems:") | |
| for i, pattern in enumerate(relevant_patterns[:3]): # Top 3 | |
| if pattern["type"] == "rule": | |
| enhancements.append(f" {i+1}. {pattern['key']}: {pattern['value']}") | |
| elif pattern["type"] == "morpheme_segmentation": | |
| val = pattern["value"] | |
| enhancements.append(f" {i+1}. Morpheme pattern: {val[0]} = {val[1]} + {val[2]}") | |
| enhancements.append("Consider these patterns in your analysis.\n") | |
| return base_prompt + "\n" + "\n".join(enhancements) if enhancements else base_prompt | |
| def rl_training_loop(self, df: pd.DataFrame, n_iterations: int = 2, samples_per_iter: int = 3, | |
| time_cap_seconds: float = None) -> dict: | |
| """ | |
| Main RL training loop: sample, evaluate, reinforce, repeat. | |
| Returns optimized hyperparameters and learned patterns. | |
| FIX: previously the only time check inside the loop happened once per | |
| sample row (before starting a config's generation batch), so a single | |
| slow batch of generations could still blow well past the intended | |
| budget and starve the main per-row loop of time, resulting in a | |
| submission full of empty timeout-fallback predictions. Now there is | |
| a hard wall-clock deadline (time_cap_seconds, relative to START) that | |
| is checked before every single model.generate() call, not just once | |
| per sample/config. | |
| """ | |
| print(f"\n=== REINFORCEMENT LEARNING ({n_iterations} iterations) ===", flush=True) | |
| if time_cap_seconds is None: | |
| time_cap_seconds = TIME_LIMIT * 0.15 # hard default cap: 15% of total budget | |
| rl_deadline = START + time_cap_seconds | |
| best_config = None | |
| best_avg_reward = 0.0 | |
| # Configurations to try (will be refined based on rewards) | |
| configs_to_try = [ | |
| {"num_paths": 3, "temperature": 0.2, "top_p": 0.9}, | |
| {"num_paths": 4, "temperature": 0.3, "top_p": 0.95}, | |
| {"num_paths": 5, "temperature": 0.25, "top_p": 0.92}, | |
| ] | |
| for iteration in range(n_iterations): | |
| if time.time() > rl_deadline: | |
| print(f" RL time cap reached before iteration {iteration+1}, stopping", flush=True) | |
| break | |
| print(f"\n--- RL Iteration {iteration + 1}/{n_iterations} ---", flush=True) | |
| # Sample test data | |
| sample_df = self.sample_test_data(df, samples_per_iter) | |
| print(f"Sampled {len(sample_df)} items for training", flush=True) | |
| iteration_rewards = [] | |
| for _, row in sample_df.iterrows(): | |
| if time.time() > rl_deadline: | |
| print(" RL time cap reached mid-iteration, stopping", flush=True) | |
| break | |
| task_type = row.get("task_type", "translation") | |
| eval_type = row.get("eval_type", "chr_f1") | |
| eval_metric = "exact" if eval_type.startswith("exact") else "chrF" | |
| k = count_items(row["query"]) | |
| # Build base prompt | |
| base_prompt = build_router_aware_prompt( | |
| row["context"], row["query"], task_type, eval_type, k | |
| ) | |
| # Enhance with learned patterns | |
| enhanced_prompt = self.build_rl_enhanced_prompt(base_prompt, task_type) | |
| # Try different configurations | |
| for config in configs_to_try[:2]: # Try first 2 configs | |
| if time.time() > rl_deadline: | |
| print(" RL time cap reached before config trial, stopping", flush=True) | |
| break | |
| start_t = time.time() | |
| # Generate paths | |
| messages = [ | |
| {"role": "system", "content": BASE_SYSTEM_PROMPT}, | |
| {"role": "user", "content": enhanced_prompt}, | |
| ] | |
| model_inputs = tok.apply_chat_template( | |
| messages, add_generation_prompt=True, return_tensors="pt" | |
| ) | |
| ids = model_inputs['input_ids'].to(model.device) | |
| # Also respect the absolute competition time limit | |
| if time_left() < TIME_LIMIT * 0.4: | |
| print(f" RL sample aborted: critical time ({time_left():.0f}s)", flush=True) | |
| break | |
| paths = [] | |
| for _ in range(config["num_paths"]): | |
| # FIX: check the deadline before every single generation, | |
| # not just once per sample/config. | |
| if time.time() > rl_deadline or time_left() < TIME_LIMIT * 0.35: | |
| print(" RL time cap reached mid-generation, stopping this batch", flush=True) | |
| break | |
| try: | |
| with torch.no_grad(): | |
| out = model.generate( | |
| ids, | |
| max_new_tokens=MAX_NEW_TOKENS, | |
| do_sample=True, | |
| temperature=config["temperature"], | |
| top_p=config["top_p"], | |
| pad_token_id=tok.eos_token_id, | |
| ) | |
| text = tok.decode(out[0][ids.shape[-1]:], skip_special_tokens=True).strip() | |
| paths.append(text) | |
| except Exception: | |
| paths.append("") | |
| gen_time = time.time() - start_t | |
| if not paths: | |
| continue | |
| # Extract answers and compute consensus | |
| path_answers = [] | |
| for path in paths: | |
| if path: | |
| ans = extract_answers_from_text(path, k) | |
| path_answers.append(enforce_k_length(ans, k)) | |
| if path_answers: | |
| consensus, confidence, _ = consensus_vote(path_answers, eval_metric) | |
| consensus = enforce_k_length(consensus, k) | |
| # Compute reward | |
| reward = self.compute_reward( | |
| paths, consensus, confidence, k, gen_time | |
| ) | |
| iteration_rewards.append(reward["total"]) | |
| # Extract and reinforce patterns if reward is good | |
| if reward["total"] > 0.5: | |
| patterns = self.extract_patterns(paths, consensus, task_type) | |
| self.update_learned_patterns(patterns, reward["total"]) | |
| # Track best config | |
| if reward["total"] > best_avg_reward: | |
| best_avg_reward = reward["total"] | |
| best_config = config | |
| avg_reward = sum(iteration_rewards) / len(iteration_rewards) if iteration_rewards else 0.0 | |
| self.iteration_rewards.append(avg_reward) | |
| print(f" Average reward: {avg_reward:.3f}", flush=True) | |
| print(f" Learned patterns: {len(self.learned_patterns)}", flush=True) | |
| # Determine final config | |
| if best_config is None: | |
| best_config = configs_to_try[0] # Default | |
| print(f"\n=== RL TRAINING COMPLETE ===", flush=True) | |
| print(f"Best config: N={best_config['num_paths']}, T={best_config['temperature']}", flush=True) | |
| print(f"Best reward: {best_avg_reward:.3f}", flush=True) | |
| print(f"Total learned patterns: {len(self.learned_patterns)}", flush=True) | |
| print(f"Time spent in RL: {time.time() - START:.1f}s (cap was {time_cap_seconds:.0f}s)", flush=True) | |
| return { | |
| "best_config": best_config, | |
| "learned_patterns": self.learned_patterns, | |
| "reward_history": self.iteration_rewards, | |
| } | |
| def get_learned_pattern_summary(self) -> str: | |
| """Return summary of learned patterns for explanation.""" | |
| if not self.learned_patterns: | |
| return "No patterns learned yet" | |
| top_patterns = sorted( | |
| self.learned_patterns.values(), | |
| key=lambda x: x["success_score"], | |
| reverse=True | |
| )[:5] | |
| summary = [] | |
| for p in top_patterns: | |
| summary.append(f" - {p['key']} (score: {p['success_score']:.2f}, n={p['count']})") | |
| return "\n".join(summary) | |
| # Global RL learner instance | |
| rl_learner = ReinforcementLearner() | |
| # ============================================================================= | |
| # 4. DETERMINISTIC ALIGNMENT & FALLBACK GUARDRAIL | |
| # ============================================================================= | |
| def enforce_k_length(answers: list, k: int) -> list: | |
| """Ensure output is exactly K items.""" | |
| answers = list(answers)[:k] | |
| while len(answers) < k: | |
| answers.append("") | |
| return answers | |
| def chrF_fallback_recovery(paths: list, expected_k: int, target_item_idx: int) -> str: | |
| """When consensus fails, try chrF-based selection.""" | |
| candidates = [] | |
| for path in paths: | |
| answers = extract_answers_from_text(path, expected_k) | |
| if target_item_idx < len(answers) and answers[target_item_idx]: | |
| candidates.append(answers[target_item_idx]) | |
| if not candidates: | |
| return "" | |
| counts = Counter(candidates) | |
| return counts.most_common(1)[0][0] | |
| def process_row_with_architecture(row: pd.Series, use_rl_prompt: bool = True) -> dict: | |
| """Main processing function with optional RL enhancement.""" | |
| task_type = row.get("task_type", "translation") | |
| eval_type = row.get("eval_type", "chr_f1") | |
| context = row["context"] | |
| query = row["query"] | |
| k = count_items(query) | |
| profile = get_task_profile(task_type) | |
| eval_metric = "exact" if eval_type.startswith("exact") else "chrF" | |
| print(f"\n[{row['id']}] Task={task_type}, Eval={eval_metric}, K={k}", flush=True) | |
| base_prompt = build_router_aware_prompt(context, query, task_type, eval_type, k) | |
| # Enhance with learned patterns if RL is active | |
| if use_rl_prompt: | |
| prompt = rl_learner.build_rl_enhanced_prompt(base_prompt, task_type) | |
| else: | |
| prompt = base_prompt | |
| print(f" Generating N={NUM_PATHS} paths with Few-Shot Rules...", flush=True) | |
| paths = generate_n_paths(prompt, n=NUM_PATHS, temperature=TEMPERATURE, top_p=TOP_P) | |
| path_answers = [] | |
| all_rules = [] | |
| for i, path in enumerate(paths): | |
| if path: | |
| answers = extract_answers_from_text(path, k) | |
| answers = enforce_k_length(answers, k) | |
| path_answers.append(answers) | |
| rules = extract_rules(path) | |
| all_rules.extend(rules) | |
| print(f" Path {i+1}: {answers}", flush=True) | |
| if path_answers: | |
| consensus, confidence, voting_explanation = consensus_vote(path_answers, eval_metric) | |
| consensus = enforce_k_length(consensus, k) | |
| print(f" Consensus: {consensus} (confidence={confidence:.2f})", flush=True) | |
| else: | |
| consensus = [""] * k | |
| confidence = 0.0 | |
| voting_explanation = "No valid paths" | |
| final_answers = enforce_k_length(consensus, k) | |
| if confidence < 0.5 and path_answers: | |
| for i in range(k): | |
| if not final_answers[i] or not final_answers[i].strip(): | |
| recovered = chrF_fallback_recovery(paths, k, i) | |
| if recovered: | |
| final_answers[i] = recovered | |
| print(f" chrF recovery for item {i+1}: {recovered}", flush=True) | |
| final_answers = enforce_k_length(final_answers, k) | |
| best_explanation = "" | |
| for path in paths: | |
| if path: | |
| best_explanation = extract_explanation(path) | |
| if best_explanation: | |
| break | |
| unique_rules = list(dict.fromkeys([r for r in all_rules if r])) | |
| rules_summary = " | ".join(unique_rules[:5]) if unique_rules else "No rules" | |
| full_explanation = f"{voting_explanation} | {best_explanation[:350]} | Rules: {rules_summary[:250]}" | |
| return { | |
| "id": row["id"], | |
| "pred": json.dumps(final_answers, ensure_ascii=False), | |
| "explanation": full_explanation[:1200], | |
| } | |
| # ============================================================================= | |
| # SETUP: Reinforcement Learning + Hyperparameter Tuning | |
| # | |
| # FIX: both phases are now hard-capped in wall-clock time and gated by much | |
| # more conservative time-remaining thresholds, so neither can meaningfully | |
| # eat into the budget needed for the main per-row processing loop. RL is | |
| # capped at 15% of TIME_LIMIT; HP tuning only runs if RL did not run/found | |
| # nothing, and internally bails once time_left() drops below 60% of budget. | |
| # ============================================================================= | |
| rl_results = None | |
| # Phase 1: Reinforcement Learning from test samples (if time permits and not disabled) | |
| if not DISABLE_RL and time_left() > TIME_LIMIT * 0.85: # only with very large headroom | |
| try: | |
| print("\n" + "="*60, flush=True) | |
| print("PHASE 1: REINFORCEMENT LEARNING FROM TEST DATA", flush=True) | |
| print("="*60, flush=True) | |
| rl_results = rl_learner.rl_training_loop( | |
| df, | |
| n_iterations=2, | |
| samples_per_iter=3, | |
| time_cap_seconds=TIME_LIMIT * 0.15, | |
| ) | |
| # Apply RL-discovered config | |
| if rl_results and rl_results.get("best_config"): | |
| cfg = rl_results["best_config"] | |
| NUM_PATHS = cfg["num_paths"] | |
| TEMPERATURE = cfg["temperature"] | |
| TOP_P = cfg.get("top_p", TOP_P) | |
| print(f"\n>>> Applied RL config: N={NUM_PATHS}, T={TEMPERATURE}, top_p={TOP_P}", flush=True) | |
| print(f"\n>>> Learned patterns:") | |
| print(rl_learner.get_learned_pattern_summary(), flush=True) | |
| except Exception as e: | |
| print(f"RL training error: {e}", flush=True) | |
| # Phase 2: Traditional hyperparameter tuning (if RL didn't run or failed) | |
| if not rl_results and time_left() > TIME_LIMIT * 0.7: | |
| try: | |
| print("\n" + "="*60, flush=True) | |
| print("PHASE 2: HYPERPARAMETER TUNING", flush=True) | |
| print("="*60, flush=True) | |
| best_config = grid_search_hyperparams(VALIDATION_EXAMPLES, max_configs=4) | |
| apply_tuned_hyperparams(best_config) | |
| except Exception as e: | |
| print(f"HP tuning error: {e}", flush=True) | |
| apply_tuned_hyperparams(None) | |
| else: | |
| if not rl_results: | |
| print("Skipping tuning - time constrained", flush=True) | |
| apply_tuned_hyperparams(None) | |
| print(f"\n>>> Time remaining before main loop: {time_left():.0f}s / {TIME_LIMIT}s", flush=True) | |
| # ============================================================================= | |
| # MAIN PROCESSING LOOP (with RL-enhanced prompts) | |
| # ============================================================================= | |
| rows_out = [] | |
| n_rows = len(df) | |
| # Main processing with time awareness and graceful degradation | |
| for i, r in df.iterrows(): | |
| expected_k = count_items(r["query"]) | |
| remaining_rows = n_rows - i | |
| per_row_budget = (time_left() - SAFETY_BUFFER) / max(1, remaining_rows) | |
| # Safety: absolute time cutoff | |
| if time_left() < SAFETY_BUFFER: | |
| print(f"TIMEOUT: Only {time_left():.0f}s left, using fast fallback", flush=True) | |
| rows_out.append({ | |
| "id": r["id"], | |
| "pred": json.dumps([""] * expected_k, ensure_ascii=False), | |
| "explanation": "Timeout: processing stopped", | |
| }) | |
| continue | |
| if time_left() < SAFETY_BUFFER or per_row_budget < 10: | |
| rows_out.append({ | |
| "id": r["id"], | |
| "pred": json.dumps([""] * expected_k, ensure_ascii=False), | |
| "explanation": "Timeout fallback", | |
| }) | |
| continue | |
| try: | |
| result = process_row_with_architecture(r) | |
| rows_out.append(result) | |
| except Exception as e: | |
| print(f"Error processing {r['id']}: {e}", flush=True) | |
| rows_out.append({ | |
| "id": r["id"], | |
| "pred": json.dumps([""] * expected_k, ensure_ascii=False), | |
| "explanation": f"Error: {str(e)[:150]}", | |
| }) | |
| pd.DataFrame(rows_out).to_csv("submission.csv", index=False) | |
| print(f"Progress: {len(rows_out)}/{n_rows} rows, {time_left():.0f}s left", flush=True) | |
| pd.DataFrame(rows_out, columns=["id", "pred", "explanation"]).to_csv( | |
| "submission.csv", index=False) | |
| print("wrote submission.csv", flush=True) | |
| print(f"Total time: {time.time() - START:.1f}s", flush=True) |