#!/usr/bin/env python3 """ Gaussian GRPO Training for Abstract Tokens """ import argparse import json import torch import glob import os import sys import numpy as np from pathlib import Path from tqdm import tqdm from torch.optim import AdamW from abstract_model import AbstractModel def extract_bracketed_answer(text): """Extract answer from [FINAL ANSWER: X] format.""" import re match = re.search(r'\[FINAL ANSWER:\s*(.*?)\]', text, re.IGNORECASE) if match: return match.group(1).strip() return None def normalize_answer(s): """Normalize answer for robust comparison.""" import re import string s = str(s).strip().lower() s = re.sub(r'\\$\\$.*?$\\`', '', s) s = re.sub(r'\\$', '', s) s = re.sub(r'\\text\{(.*?)\}', r'\1', s) s = s.translate(str.maketrans('', '', string.punctuation)) return ' '.join(s.split()) def compute_reward(generated_text, reference_answer, mode_sequence): """ Compute composite reward: Accuracy + Structure """ bracketed = extract_bracketed_answer(generated_text) gen_to_compare = bracketed if bracketed else generated_text gen_norm = normalize_answer(gen_to_compare) ref_norm = normalize_answer(reference_answer) if not ref_norm: answer_score = 0.0 elif gen_norm == ref_norm: answer_score = 1.0 elif ref_norm in gen_norm.split(): answer_score = 1.0 else: # Partial overlap gen_words = set(gen_norm.split()) ref_words = set(ref_norm.split()) if len(ref_words) > 0: answer_score = len(gen_words & ref_words) / len(ref_words) else: answer_score = 0.0 # Structure Reward: Did it use ? has_transition = 'T' in mode_sequence structure_score = 1.0 if has_transition else 0.0 total_reward = (0.7 * answer_score) + (0.3 * structure_score) # Boost for perfect result if answer_score == 1.0 and has_transition: total_reward = 1.0 return total_reward SYSTEM_PROMPTS = { 'boolean_expressions': "Provide your final answer (True or False) at the end of your response in this exact format: [FINAL ANSWER: X].", 'dyck_language': "Provide the completion sequence at the end of your response in this exact format: [FINAL ANSWER: X].", 'causal_judgement': "Provide your answer (Yes or No) at the end of your response in this exact format: [FINAL ANSWER: X].", 'formal_fallacies': "Provide your answer at the end of your response in this exact format: [FINAL ANSWER: X].", 'logical_deduction_three_objects': "Provide your final answer at the end of your response in this exact format: [FINAL ANSWER: X].", 'math_level1': "Provide your final numerical answer at the end of your response in this exact format: [FINAL ANSWER: X].", 'prontoqa': "Provide your answer (True or False) at the end of your response in this exact format: [FINAL ANSWER: X].", 'temporal_sequences': "Provide the next element(s) in the sequence at the end of your response in this exact format: [FINAL ANSWER: X].", 'tracking_shuffled_objects_three_objects': "Provide your answer at the end of your response in this exact format: [FINAL ANSWER: X].", 'web_of_lies': "Provide your answer (True or False) at the end of your response in this exact format: [FINAL ANSWER: X].", } def load_rl_data(data_dir, max_samples=None): all_data = [] files = glob.glob(os.path.join(data_dir, "*.jsonl")) print(f"Scanning {data_dir}...") if not files: print(f"CRITICAL WARNING: No .jsonl files found in {data_dir}") return [] print(f"Found {len(files)} files.") for f_path in files: filename = os.path.basename(f_path).replace('.jsonl', '') system_prompt = SYSTEM_PROMPTS.get(filename, None) # Fuzzy match system prompt if system_prompt is None: filename_alt = filename.replace('_', '') for key in SYSTEM_PROMPTS: if key.replace('_', '') == filename_alt: system_prompt = SYSTEM_PROMPTS[key] break try: with open(f_path, 'r') as f: for line in f: try: item = json.loads(line) if 'prompt' in item and 'answer' in item: item['system_prompt'] = system_prompt all_data.append(item) except: continue except Exception as e: print(f"Error reading {f_path}: {e}") if max_samples: import random random.shuffle(all_data) all_data = all_data[:max_samples] print(f"Loaded {len(all_data)} valid training samples.") return all_data def train(args): device = 'cuda:0' if torch.cuda.is_available() else 'cpu' print(f"Device: {device}") train_data = load_rl_data(args.data_dir, args.max_samples) if not train_data: print("ERROR: Training data is empty. Exiting.") sys.exit(1) print(f"Loading Abstract model from {args.abstract_model}...") model = AbstractModel.load_from_directory(args.abstract_model, args.sft_model, device=device) try: print("Compiling model backbone with torch.compile...") model.model_backbone = torch.compile(model.model_backbone) except Exception as e: print(f"Warning: Could not compile model: {e}") model.set_trainable_params('abstract') model.train() optimizer = AdamW(filter(lambda p: p.requires_grad, model.parameters()), lr=args.lr) global_step = 0 for epoch in range(args.epochs): print(f"Epoch {epoch+1}/{args.epochs}") np.random.shuffle(train_data) progress = tqdm(range(0, len(train_data), args.batch_size)) for i in progress: batch = train_data[i : i + args.batch_size] batch_loss = 0.0 optimizer.zero_grad() for item in batch: try: prompt = item['prompt'] reference = item['answer'] sys_prompt = item.get('system_prompt', None) messages = [] if sys_prompt: messages.append({"role": "system", "content": sys_prompt}) messages.append({"role": "user", "content": prompt}) formatted_prompt = model.tokenizer.apply_chat_template( messages, tokenize=False, add_generation_prompt=True ) input_ids = model.tokenizer( formatted_prompt, return_tensors='pt', add_special_tokens=False )['input_ids'].to(model.device).squeeze(0) group_rewards = [] group_log_probs = [] for gen_idx in range(args.group_size): result = model.forward( input_ids, max_length=args.max_length, temperature=args.temperature, sigma=args.sigma, sample=True, no_grad=False ) gen_ids = result['generated_tokens'].tolist() gen_text = model.tokenizer.decode(gen_ids, skip_special_tokens=True) r = compute_reward(gen_text, reference, result['mode_sequence']) group_rewards.append(r) if len(result['log_probs']) > 0: group_log_probs.append(result['log_probs'].sum()) else: group_log_probs.append(torch.tensor(0.0, device=model.device, requires_grad=True)) rewards_np = np.array(group_rewards) mean_r = rewards_np.mean() std_r = rewards_np.std() + 1e-8 advantages = (rewards_np - mean_r) / std_r prompt_loss = 0.0 valid_items = 0 for adv, log_prob_sum in zip(advantages, group_log_probs): if log_prob_sum.requires_grad: adv_tensor = torch.tensor(adv, device=model.device, dtype=log_prob_sum.dtype) prompt_loss += -1.0 * (adv_tensor * log_prob_sum) valid_items += 1 if valid_items > 0: prompt_loss = prompt_loss / valid_items prompt_loss.backward() batch_loss += prompt_loss.item() except Exception as e: print(f"Error in batch: {e}") continue torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) optimizer.step() global_step += 1 optimizer.zero_grad() torch.cuda.empty_cache() if global_step % args.save_steps == 0: save_path = os.path.join(args.output, f"step_{global_step}") model.save_to_directory(save_path) model.save_to_directory(os.path.join(args.output, "final")) print("Done.") if __name__ == "__main__": torch.set_float32_matmul_precision('high') parser = argparse.ArgumentParser() parser.add_argument("--sft-model", required=True) parser.add_argument("--abstract-model", required=True) parser.add_argument("--data-dir", required=True) parser.add_argument("--output", required=True) parser.add_argument("--group-size", type=int, default=4) parser.add_argument("--batch-size", type=int, default=1) parser.add_argument("--lr", type=float, default=1e-5) parser.add_argument("--epochs", type=int, default=1) parser.add_argument("--max-length", type=int, default=256) parser.add_argument("--temperature", type=float, default=1.0) parser.add_argument("--sigma", type=float, default=0.1) parser.add_argument("--max-samples", type=int, default=None) parser.add_argument("--save-steps", type=int, default=50) args = parser.parse_args() os.makedirs(args.output, exist_ok=True) train(args)