FST_code / src /lmr /generation /generation_grid.py
jasonfan's picture
2026-03-19
3b2d368 verified
Raw
History Blame Contribute Delete
35.3 kB
#!/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