#!/usr/bin/env python3 """ Unified Generator with TinyGSM/GSM8K Support + Grid Search (Defaults requested) - Uses function name `simple_math_problem` - Real-time accuracy display during generation - 5-fold evaluation support (config.five_fold = True) - Grid search over (top_p, temperature) with caching per combination. - Per-combination timeout using SIGALRM on Unix. On non-Unix platforms we record timeout but cannot force-stop model.generate. - DEFAULTS (as requested): * Folds to test by default: fold 0 and fold 1 (i.e. two folds) * top_p grid: [0.7, 0.8, 0.9] * temperature grid: [0.2, 0.5, 0.7, 0.8] """ import re import json import time import ast import signal import traceback from pathlib import Path from typing import List, Dict, Optional, Union import torch from tqdm import tqdm # Default prompts DEFAULT_CHAT_PROMPT = """[USER] Hello, how are you today? [SYSTEM] I am""" class Generator: """ 统一生成器,支持: - chat: 标准对话生成 - tinygsm dataset: 自动在 GSM8K test 上评估(执行生成的 Python 代码) 支持五折评估:将数据拆成 5 份,分别作为 test fold 执行评估并保存每个 fold 的结果 支持 grid search:测试不同 top_p / temperature 的组合,并缓存结果(跳过已存在文件) 默认行为:在 tinygsm 模式下默认执行 run_grid_evaluations(fold_indices=[0,1]) """ def __init__(self, config, model, tokenizer, device="cuda", output_dir=None): # the generation-related config is expected under config.generation; fallback to config if not present self.config = config.generation if hasattr(config, "generation") else config # but keep a reference to full config for dataset name etc. self.full_config = config self.model = model.eval().to(device) self.tokenizer = tokenizer self.device = device # 将 output_dir 保存为实例属性,供保存结果使用 self.output_dir = Path(output_dir or getattr(self.config, 'output_dir', 'output')) # 检测 dataset 类型 self.dataset_name = getattr(self.full_config, 'dataset', None) if hasattr(self.dataset_name, 'dataset_name'): self.dataset_name = self.dataset_name.dataset_name self.is_tinygsm = self.dataset_name == "tinygsm" print(f"[Generator] Initialized") print(f" Dataset: {self.dataset_name}") print(f" Device: {device}") if self.is_tinygsm: print(f" Mode: TinyGSM - Will evaluate on GSM8K test set with CODE EXECUTION") print(f" Five-fold enabled: {getattr(self.config, 'five_fold', True)}") # -------------------- Public generation entry -------------------- def generate(self, prompt: Optional[Union[str, List[str]]] = None): """ 主生成入口 - 如果是 tinygsm dataset: 自动加载 GSM8K test 并评估(默认调用 run_grid_evaluations on folds 0 & 1) - 否则: 标准 chat 生成 """ if self.is_tinygsm: return self._generate_tinygsm() else: return self._generate_chat(prompt) # -------------------- Chat -------------------- def _generate_chat(self, prompt: Optional[str] = None): """标准 chat 模式生成""" if prompt is None: prompt_text = getattr(self.config, "prompt", DEFAULT_CHAT_PROMPT) else: prompt_text = prompt max_len = getattr(self.config, "max_new_tokens", 512) temperature = getattr(self.config, "temperature", 0.7) top_p = getattr(self.config, "top_p", 0.9) return_gen_only = getattr(self.config, "return_generation_only", True) input_ids = torch.tensor( [self.tokenizer.encode(prompt_text)], dtype=torch.long, device=self.device ) with torch.no_grad(): outputs = self.model.generate( input_ids, max_generation_length=max_len, tokenizer=self.tokenizer, temperature=temperature, top_p=top_p, return_generation_only=return_gen_only ) generated_texts = [] for seq in outputs: text = self.tokenizer.decode(seq.tolist()) generated_texts.append(text) if getattr(self.config, "save_to_file", False): self._save_generations(generated_texts, "generations.txt") for text in generated_texts: print(f"{prompt_text} [{text}]") return generated_texts # -------------------- TinyGSM / GSM8K -------------------- def _generate_tinygsm(self): """TinyGSM mode: 在 GSM8K test set 上生成并验证 DEFAULT: run grid evaluations on folds [0,1] with default grids """ print("\n" + "="*60) print("TinyGSM Dataset Detected") print("Evaluating on GSM8K Test Set") print("="*60 + "\n") print("Step 1: Loading GSM8K Test Set...") gsm8k_data = self._load_gsm8k_test() if not gsm8k_data or not gsm8k_data['questions']: print("❌ Failed to load GSM8K test set") return None questions = gsm8k_data['questions'] ground_truth = gsm8k_data['answers'] max_samples = getattr(self.config, 'max_samples', None) if max_samples is not None and len(questions) > max_samples: print(f"Limiting to {max_samples} samples") questions = questions[:max_samples] ground_truth = ground_truth[:max_samples] print(f"✓ Loaded {len(questions)} questions\n") # DEFAULT action: run grid evaluations on folds 0 and 1 default_topps = getattr(self.config, 'grid_topps', [0.7, 0.8, 0.9]) default_temps = getattr(self.config, 'grid_temperatures', [0.2, 0.5, 0.7, 0.8]) timeout_seconds = int(getattr(self.config, 'combo_timeout_seconds', 180)) print("[Generator] Default action for tinygsm: running run_grid_evaluations on folds [0,1].") return self.run_grid_evaluations( temps=default_temps, topps=default_topps, n_folds=5, only_fold0=False, fold_indices=[0, 1], ) def _load_gsm8k_test(self) -> Dict[str, List]: """加载 GSM8K test set""" try: from datasets import load_dataset print("Loading from HuggingFace...") ds = load_dataset("openai/gsm8k", "main", split="test") questions = [] answers = [] with tqdm(total=len(ds), desc="Loading", unit="sample") as pbar: for sample in ds: q = (sample.get("question", "") or "").strip() a = (sample.get("answer", "") or "").strip() if not q or not a: pbar.update(1) continue answer_num = self._extract_gsm8k_answer(a) questions.append(q) answers.append(answer_num) pbar.update(1) return {'questions': questions, 'answers': answers} except Exception as e: print(f"❌ Error loading GSM8K: {e}") return {'questions': [], 'answers': []} def _extract_gsm8k_answer(self, answer_text: str) -> str: """从 GSM8K 答案中提取数值""" match = re.search(r'####\s*([+-]?[\d,]+\.?\d*)', answer_text) if match: return match.group(1).replace(',', '').strip() match = re.search(r'result\s*=\s*([+-]?[\d,]+\.?\d*)', answer_text, re.IGNORECASE) if match: return match.group(1).replace(',', '').strip() numbers = re.findall(r'([+-]?[\d,]+\.?\d*)', answer_text) if numbers: return numbers[-1].replace(',', '').strip() return "0" def _create_math_prompt(self, question: str) -> str: """创建结构化的数学问题提示""" return f"""<|bos|>Question: {question} Solution: def simple_math_problem() -> int: # {question} """ # -------------------- 生成并执行解决方案 -------------------- def _generate_solutions(self, questions: List[str], ground_truth: List[str]) -> Dict: """ 为每个问题生成代码解决方案 在循环中加入实时准确率计算和进度条更新 """ max_len = getattr(self.config, 'max_new_tokens', 512) temperature = getattr(self.config, 'temperature', 0.7) top_p = getattr(self.config, 'top_p', 0.9) results = { 'questions': [], 'prompts': [], 'generated_solutions': [], 'extracted_code': [], 'executed_answers': [], 'full_outputs': [] } correct_count = 0 total_processed = 0 # 将问题和标准答案打包,并在 tqdm 中遍历 pbar = tqdm(zip(questions, ground_truth), total=len(questions), desc="Generating", unit="q") for idx, (question, truth_val) in enumerate(pbar): prompt = self._create_math_prompt(question) input_ids = torch.tensor( [self.tokenizer.encode(prompt)], dtype=torch.long, device=self.device ) with torch.no_grad(): outputs = self.model.generate( input_ids, max_generation_length=max_len, tokenizer=self.tokenizer, temperature=temperature, top_p=top_p, return_generation_only=False ) full_text = self.tokenizer.decode(outputs[0].tolist()) # 调试打印 (减少频率) if idx == 0: print(f"\n{'='*60}") print(f"🔍 DEBUG Sample {idx + 1}: Raw generated text (first 500 chars):") print(f"{'='*60}") print(full_text[:500]) print(f"{'='*60}\n") if full_text.startswith(prompt): solution = full_text[len(prompt):] else: solution = full_text solution = solution.replace("<|eos|>", "").replace("<|bos|>", "").strip() code = self._extract_function_code(full_text) if code: answer = self._execute_math_code(code) else: answer = None # --- 实时验证和进度条更新 --- is_correct = False if truth_val is not None: is_correct = self._compare_answers(answer, truth_val) if is_correct: correct_count += 1 total_processed += 1 # 更新进度条后缀显示准确率 current_acc = correct_count / total_processed if total_processed > 0 else 0.0 pbar.set_postfix({"Acc": f"{current_acc:.2%}"}) # ------------------------------ results['questions'].append(question) results['prompts'].append(prompt) results['generated_solutions'].append(solution) results['extracted_code'].append(code if code else "No code extracted") results['executed_answers'].append(answer) results['full_outputs'].append({ 'question': question, 'solution': solution, 'code': code, 'answer': answer }) if idx < 3 or idx == len(questions) - 1: print(f"\n{'─'*60}") print(f"Sample {idx + 1}") print(f"{'─'*60}") print(f"Q: {question[:100]}...") if code: print(f"Code:\n{code[:200]}...") print(f"Answer: {answer if answer is not None else '[None]'}") if truth_val: status = "✓" if is_correct else "✗" print(f"Truth: {truth_val} {status}") print(f"{'─'*60}") return results # -------------------- 五折支持 -------------------- def _split_into_folds(self, questions: List[str], answers: List[str], n_folds: int = 5): """按顺序把 questions/answers 拆成 n_folds 份,尽量平均""" assert len(questions) == len(answers) total = len(questions) folds_q = [] folds_a = [] base = total // n_folds remainder = total % n_folds start = 0 for i in range(n_folds): size = base + (1 if i < remainder else 0) end = start + size folds_q.append(questions[start:end]) folds_a.append(answers[start:end]) start = end return folds_q, folds_a def _generate_tinygsm_five_fold(self, questions: List[str], ground_truth: List[str], n_folds: int = 5): """ 五折评估(备用): - 将数据分成 n_folds 份 - 对每个 fold 单独作为测试集,调用 _generate_solutions - 保存每个 fold 的结果到 self.output_dir/gsm8k_results_fold_{i+1}.txt - 返回包含每个 fold 详细结果和总体汇总的字典 """ print("Running 5-fold evaluation...") folds_q, folds_a = self._split_into_folds(questions, ground_truth, n_folds=n_folds) all_folds_results = [] fold_accuracies = [] self.output_dir.mkdir(parents=True, exist_ok=True) for i in range(n_folds): print(f"\n{'#'*40}\nEvaluating Fold {i+1}/{n_folds} - test size: {len(folds_q[i])}\n{'#'*40}\n") # 对当前 fold 执行生成与验证 fold_results = self._generate_solutions(folds_q[i], folds_a[i]) # 验证(_verify_answers 返回的 total/accuracy 等) fold_verification = self._verify_answers(fold_results, folds_a[i]) fold_results['verification'] = fold_verification fold_results['ground_truth'] = folds_a[i] # 保存 fold 文件 fold_out_path = self.output_dir / f"gsm8k_results_fold_{i+1}.txt" with open(fold_out_path, 'w', encoding='utf-8') as f: for j, item in enumerate(fold_results['full_outputs'], 1): f.write(f"{'='*60}\n") f.write(f"Fold {i+1} - Problem {j}\n") f.write(f"{'='*60}\n") f.write(f"Question:\n{item['question']}\n\n") f.write(f"Solution:\n{item['solution']}\n\n") if item.get('code'): f.write(f"Extracted Code:\n{item['code']}\n\n") f.write(f"Answer: {item['answer']}\n") if 'ground_truth' in fold_results and j <= len(fold_results['ground_truth']): truth = fold_results['ground_truth'][j-1] match = "✓" if self._compare_answers(item['answer'], truth) else "✗" f.write(f"Truth: {truth} {match}\n") f.write(f"{'='*60}\n\n") print(f"[Generator] Saved fold results to {fold_out_path}") # 保存 verification 简要文件 verify_path = self.output_dir / f"verification_fold_{i+1}.txt" with open(verify_path, 'w', encoding='utf-8') as f: v = fold_verification f.write(f"{'='*60}\n") f.write(f"GSM8K Verification Results - Fold {i+1}\n") f.write(f"{'='*60}\n") f.write(f"Total: {v['total']}\n") f.write(f"Correct: {v['correct']}\n") f.write(f"Incorrect: {v['incorrect']}\n") f.write(f"No Answer: {v['no_answer']}\n") f.write(f"Accuracy: {v['accuracy']*100:.2f}%\n") f.write(f"{'='*60}\n\n") print(f"[Generator] Saved fold verification to {verify_path}") all_folds_results.append(fold_results) fold_accuracies.append(fold_verification.get('accuracy', 0.0)) # 打印单折汇总 print(f"Fold {i+1} Accuracy: {fold_verification.get('accuracy', 0.0)*100:.2f}%") # 总体汇总 overall_mean_acc = sum(fold_accuracies) / len(fold_accuracies) if fold_accuracies else 0.0 summary = { 'n_folds': n_folds, 'fold_accuracies': fold_accuracies, 'mean_accuracy': overall_mean_acc, 'folds': all_folds_results } # 保存总体 summary summary_path = self.output_dir / "gsm8k_5fold_summary.txt" with open(summary_path, 'w', encoding='utf-8') as f: f.write(f"{'='*60}\n") f.write("GSM8K 5-Fold Summary\n") f.write(f"{'='*60}\n") for idx, acc in enumerate(fold_accuracies, 1): f.write(f"Fold {idx} Accuracy: {acc*100:.2f}%\n") f.write(f"\nMean Accuracy: {overall_mean_acc*100:.2f}%\n") f.write(f"{'='*60}\n") print(f"[Generator] Saved 5-fold summary to {summary_path}") # 打印总体结果 print("\n" + "="*60) print("5-Fold Evaluation Summary") print(f"Mean Accuracy: {overall_mean_acc*100:.2f}%") for idx, acc in enumerate(fold_accuracies, 1): print(f" Fold {idx}: {acc*100:.2f}%") print("="*60 + "\n") return summary # -------------------- 提取/执行/验证 -------------------- def _extract_function_code(self, generated_text: str) -> Optional[str]: """Extract and clean function code from generated text""" func_start = generated_text.find("def simple_math_problem") if func_start == -1: return None func_text = generated_text[func_start:] lines = func_text.split('\n') cleaned_lines = [] for i, line in enumerate(lines): if i == 0: cleaned_lines.append("def simple_math_problem() -> int:") else: stripped_line = line.strip() # Stop at next top-level def or common terminators if stripped_line and not line.startswith(' ') and not line.startswith('\t'): if 'def ' in line or stripped_line.startswith('Answer:') or stripped_line.startswith('```'): break if stripped_line in ['```', '```python', '```py'] or stripped_line.startswith('Answer:') or '<|eos|>' in stripped_line: break clean_line = line.rstrip() if clean_line.strip(): # skip docstrings and comments stripped = clean_line.strip() if stripped.startswith('"""') or stripped.startswith("'''") or stripped.startswith('#'): continue # ensure indentation if not clean_line.startswith(' '): clean_line = ' ' + clean_line.strip() cleaned_lines.append(clean_line) else: # preserve blank line inside function as a blank indented line cleaned_lines.append('') # remove trailing blank lines while cleaned_lines and not cleaned_lines[-1].strip(): cleaned_lines.pop() has_return = any('return' in line for line in cleaned_lines) if not has_return: cleaned_lines.append(' return result') final_code = '\n'.join(cleaned_lines) return final_code def _execute_math_code(self, code: str) -> Optional[float]: """Safely execute mathematical code""" try: if not code: return None if 'return' not in code: code += '\n return result' full_code = code + '\n\nresult = simple_math_problem()' try: ast.parse(full_code) except SyntaxError: return None safe_globals = { "__builtins__": { "abs": abs, "round": round, "min": min, "max": max, "sum": sum, "int": int, "float": float, "str": str, "len": len, "range": range, "list": list, "dict": dict, "set": set, "tuple": tuple, "True": True, "False": False, "None": None } } safe_locals = {} exec(full_code, safe_globals, safe_locals) if 'result' in safe_locals: try: result_value = float(safe_locals['result']) return result_value except Exception: return None else: return None except Exception: return None def _verify_answers(self, results: Dict, ground_truth: List[str]) -> Dict: """验证答案""" predictions = results['executed_answers'] verification = { 'total': len(ground_truth), 'correct': 0, 'incorrect': 0, 'no_answer': 0, 'accuracy': 0.0, 'details': [] } for idx, (pred, truth) in enumerate(zip(predictions, ground_truth)): if pred is None: verification['no_answer'] += 1 is_correct = False else: is_correct = self._compare_answers(pred, truth) if is_correct: verification['correct'] += 1 else: verification['incorrect'] += 1 verification['details'].append({ 'index': idx, 'question': results['questions'][idx], 'predicted': pred, 'ground_truth': truth, 'correct': is_correct }) verification['accuracy'] = ( verification['correct'] / verification['total'] if verification['total'] > 0 else 0.0 ) return verification def _compare_answers(self, pred: Union[str, float, None], truth: str) -> bool: """数值比较""" if pred is None: return False pred_str = str(pred).strip().replace(',', '') truth_str = str(truth).strip().replace(',', '') if pred_str == truth_str: return True try: pred_num = float(pred_str) truth_num = float(truth_str) return abs(pred_num - truth_num) < 1e-6 except (ValueError, TypeError): pass return False def _print_verification_summary(self, verification: Dict): """打印验证总结""" print(f"\n{'='*60}") print("GSM8K Evaluation Results") print(f"{'='*60}") print(f"Total: {verification['total']}") print(f"Correct: {verification['correct']} ({verification['accuracy']*100:.2f}%)") print(f"Incorrect: {verification['incorrect']}") print(f"No Answer: {verification['no_answer']}") print(f"{'='*60}") incorrect = [d for d in verification['details'] if not d['correct']][:3] if incorrect: print("\nSample Incorrect Predictions:") for detail in incorrect: print(f" Q: {detail['question'][:60]}...") print(f" Predicted: {detail['predicted']}, Truth: {detail['ground_truth']}\n") # -------------------- 保存 -------------------- def _save_generations(self, texts: List[str], filename: str): """保存 chat 生成结果""" output_dir = Path(getattr(self.config, 'output_dir', self.output_dir)) output_dir.mkdir(parents=True, exist_ok=True) output_path = output_dir / filename with open(output_path, 'w', encoding='utf-8') as f: for i, text in enumerate(texts, 1): f.write(f"{'='*60}\n") f.write(f"Generation {i}\n") f.write(f"{'='*60}\n") f.write(text + "\n\n") print(f"[Generator] Saved to {output_path}") def _save_tinygsm_results(self, results: Dict): """保存 TinyGSM/GSM8K 评估结果""" output_dir = Path(self.output_dir) output_dir.mkdir(parents=True, exist_ok=True) full_path = output_dir / "gsm8k_results.txt" with open(full_path, 'w', encoding='utf-8') as f: for i, item in enumerate(results['full_outputs'], 1): f.write(f"{'='*60}\n") f.write(f"Problem {i}\n") f.write(f"{'='*60}\n") f.write(f"Question:\n{item['question']}\n\n") f.write(f"Solution:\n{item['solution']}\n\n") if item.get('code'): f.write(f"Extracted Code:\n{item['code']}\n\n") f.write(f"Answer: {item['answer']}\n") if 'ground_truth' in results and i <= len(results['ground_truth']): truth = results['ground_truth'][i-1] match = "✓" if self._compare_answers(item['answer'], truth) else "✗" f.write(f"Truth: {truth} {match}\n") f.write(f"{'='*60}\n\n") print(f"[Generator] Saved results to {full_path}") if 'verification' in results: verify_path = output_dir / "verification.txt" with open(verify_path, 'w', encoding='utf-8') as f: v = results['verification'] f.write(f"{'='*60}\n") f.write("GSM8K Verification Results\n") f.write(f"{'='*60}\n") f.write(f"Total: {v['total']}\n") f.write(f"Correct: {v['correct']}\n") f.write(f"Incorrect: {v['incorrect']}\n") f.write(f"No Answer: {v['no_answer']}\n") f.write(f"Accuracy: {v['accuracy']*100:.2f}%\n") f.write(f"{'='*60}\n\n") for detail in v['details']: status = "✓" if detail['correct'] else "✗" f.write(f"{detail['index']+1}. {status} ") f.write(f"Pred: {detail['predicted']}, Truth: {detail['ground_truth']}\n") print(f"[Generator] Saved verification to {verify_path}") # -------------------- Grid Search / Cached combos / Timeout -------------------- def run_grid_evaluations(self, temps: Optional[List[float]] = None, topps: Optional[List[float]] = None, n_folds: int = 5, only_fold0: bool = True, fold_indices: Optional[List[int]] = None): """ 对一组 temperature 与 top_p 组合进行评估并缓存每个组合的结果。 参数说明: - temps: list of temperatures(None 时读取 self.config.grid_temperatures 或默认 [0.2,0.5,0.7,0.8]) - topps: list of top_p(None 时读取 self.config.grid_topps 或默认 [0.7,0.8,0.9]) - n_folds: folds used to split the set (default 5) - only_fold0: if True and fold_indices is None, use only fold0 - fold_indices: optional list of fold indices to evaluate (e.g. [0,1]); if provided, it overrides only_fold0 默认循环顺序:外圈 top_p -> 内圈 temperature """ topps = topps if topps is not None else getattr(self.config, 'grid_topps', [0.9]) temps = temps if temps is not None else getattr(self.config, 'grid_temperatures', [0.2, 0.5, 0.7, 0.8,1]) timeout_seconds = int(getattr(self.config, 'combo_timeout_seconds', 1800)) # default 3 minutes save_dir = Path(getattr(self.config, 'output_dir', self.output_dir)) save_dir.mkdir(parents=True, exist_ok=True) # load GSM8K once gsm = self._load_gsm8k_test() if not gsm or not gsm['questions']: print("❌ GSM8K load failed, aborting grid run.") return questions_all = gsm['questions'] answers_all = gsm['answers'] # create folds folds_q, folds_a = self._split_into_folds(questions_all, answers_all, n_folds=n_folds) # determine which folds to evaluate eval_sets = [] if fold_indices is not None: # sanitize indices and ignore out-of-range unique_idxs = [] for idx in fold_indices: if not isinstance(idx, int): continue if 0 <= idx < n_folds and idx not in unique_idxs: unique_idxs.append(idx) if not unique_idxs: print("[Grid] No valid fold indices provided (after sanitization). Nothing to run.") return for idx in unique_idxs: eval_sets.append({ 'name': f'fold{idx}', 'questions': folds_q[idx], 'answers': folds_a[idx] }) else: if only_fold0: eval_sets.append({ 'name': 'fold0', 'questions': folds_q[0], 'answers': folds_a[0] }) else: for i in range(n_folds): eval_sets.append({ 'name': f'fold{i}', 'questions': folds_q[i], 'answers': folds_a[i] }) total_combos = len(topps) * len(temps) * len(eval_sets) combo_idx = 0 for eval_set in eval_sets: qlist = eval_set['questions'] alist = eval_set['answers'] set_name = eval_set['name'] # OUTER LOOP: top_p (requested default ordering) for p in topps: for t in temps: combo_idx += 1 # encode t/p to integers to avoid weird filenames t_enc = int(round(t * 100)) p_enc = int(round(p * 100)) fname = save_dir / f"gsm8k_{set_name}_t{t_enc}_p{p_enc}.json" if fname.exists(): print(f"[Grid] Skipping existing result {fname.name}") continue print(f"[Grid] ({combo_idx}/{total_combos}) Evaluating {set_name} with top_p={p}, temperature={t}") start_time = time.time() timed_out = False result_obj = None # On Unix we can use SIGALRM to enforce timeout supports_sigalrm = hasattr(signal, "SIGALRM") def _alarm_handler(signum, frame): raise TimeoutError("Combo timed out (SIGALRM).") if supports_sigalrm: old_handler = signal.signal(signal.SIGALRM, _alarm_handler) signal.alarm(timeout_seconds) try: # Temporarily override generation params in self.config for this run old_temp = getattr(self.config, 'temperature', None) old_top_p = getattr(self.config, 'top_p', None) self.config.temperature = t self.config.top_p = p # call generation on the chosen eval set results = self._generate_solutions(qlist, alist) verification = self._verify_answers(results, alist) results['verification'] = verification results['ground_truth'] = alist elapsed = time.time() - start_time result_obj = { 'params': {'temperature': t, 'top_p': p, 'set': set_name}, 'elapsed_seconds': elapsed, 'verification': verification, 'num_questions': len(qlist), 'results_preview': results['full_outputs'][:5] # save light preview } # 保存完整 results 到磁盘(压缩或仅部分也行) with open(fname, 'w', encoding='utf-8') as wf: json.dump({ 'meta': result_obj['params'], 'elapsed_seconds': result_obj['elapsed_seconds'], 'verification': verification, 'full_outputs': results['full_outputs'] }, wf, ensure_ascii=False, indent=2) print(f"[Grid] Saved result to {fname} (elapsed {elapsed:.1f}s)") except TimeoutError as te: timed_out = True print(f"[Grid] TIMEOUT for top_p={p}, temp={t} on {set_name} after {timeout_seconds}s. Skipping.") # write a small timeout marker file so future runs skip with open(fname, 'w', encoding='utf-8') as wf: json.dump({ 'meta': {'temperature': t, 'top_p': p, 'set': set_name}, 'timed_out': True, 'message': str(te) }, wf, ensure_ascii=False, indent=2) except Exception as e: # 捕获并保存异常信息,避免整个 grid 中断 tb = traceback.format_exc() print(f"[Grid] ERROR for top_p={p}, temp={t} on {set_name}: {e}\n{tb}") with open(fname, 'w', encoding='utf-8') as wf: json.dump({ 'meta': {'temperature': t, 'top_p': p, 'set': set_name}, 'error': str(e), 'traceback': tb }, wf, ensure_ascii=False, indent=2) finally: # restore alarm & handlers if supports_sigalrm: signal.alarm(0) signal.signal(signal.SIGALRM, old_handler) # restore old config params if old_temp is not None: self.config.temperature = old_temp else: try: delattr(self.config, 'temperature') except Exception: pass if old_top_p is not None: self.config.top_p = old_top_p else: try: delattr(self.config, 'top_p') except Exception: pass print("[Grid] All combos processed. Cached results in:", str(save_dir)) return True # end of Generator class