"""Dataset preparation for GRPO training — Sangue e Grafi. Generates kinship-reasoning scenarios and packages them as a HuggingFace ``datasets.Dataset`` suitable for GRPO training with ``trl.GRPOTrainer``. Each example contains: - **prompt**: a kinship reasoning task with narrative, question, and instructions on the XML tool-calling format - **gold_answer**: the correct heir's full name - **will_clause**: the will clause as a JSON string (for reward computation) - **optimal_hops**: estimated minimum graph traversals needed Usage:: python -m src.training.prepare_dataset # 50-scenario smoke test python -c "from src.training.prepare_dataset import prepare_grpo_dataset; \\ ds = prepare_grpo_dataset(500); ds.save_to_disk('data/grpo_train')" """ from __future__ import annotations import json import logging from typing import Any from datasets import Dataset from src.graph.scenario_generator import generate_batch logger = logging.getLogger(__name__) # --------------------------------------------------------------------------- # Prompt template # --------------------------------------------------------------------------- _PROMPT_TEMPLATE: str = """\ You are a kinship reasoning agent. Navigate the knowledge graph to answer this question. Narrative: {narrative} Question: {question} The knowledge graph contains entities of types Person, Asset, and WillClause. Use tags for your thought process. Use tags for graph tool calls. Use name when you have the final answer.""" # --------------------------------------------------------------------------- # Public API # --------------------------------------------------------------------------- def prepare_grpo_dataset( n_scenarios: int = 500, seed: int = 10000, save_path: str | None = None, push_to_hub: bool = False, hub_repo_id: str | None = None, ) -> Dataset: """Generate and return a GRPO-ready HuggingFace dataset. Args: n_scenarios: Number of kinship scenarios to generate. seed: Base seed for reproducible scenario generation. Default is 10000 (training range). NEVER use benchmark seeds (42, 99, 137, 256, 500, 777, 1234, 2048, 3333, 4096). save_path: If provided, save the dataset to this directory via ``Dataset.save_to_disk()``. push_to_hub: If *True*, push the dataset to HuggingFace Hub. Requires ``hub_repo_id`` and a valid ``HF_TOKEN`` in the environment. hub_repo_id: Hub repository identifier, e.g. ``"your-org/sangue-e-grafi-grpo"``. Returns: A :class:`datasets.Dataset` with columns ``prompt``, ``gold_answer``, ``will_clause``, and ``optimal_hops``. """ # Guard against benchmark seed contamination BENCHMARK_SEEDS = {42, 99, 137, 256, 500, 777, 1234, 2048, 3333, 4096} for i in range(n_scenarios): s = seed + i if s in BENCHMARK_SEEDS: raise ValueError( f"Seed {s} is a frozen README benchmark seed. " f"Use seed >= 10000 for training data." ) logger.info("Generating %d kinship scenarios (seed=%d)…", n_scenarios, seed) bundles: list[dict[str, Any]] = generate_batch(n_scenarios, seed=seed) prompts: list[str] = [] gold_answers: list[str] = [] will_clauses: list[str] = [] optimal_hops_list: list[int] = [] for bundle in bundles: scenario = bundle["scenario"] # Format the prompt prompt = _PROMPT_TEMPLATE.format( narrative=bundle["narrative"], question=bundle["question"], ) prompts.append(prompt) # Gold answer (resolved to full name by generate_batch) gold_answers.append(bundle["gold_answer"]) # Will clause as JSON for reward-function consumption will_clause_json = json.dumps(scenario["will"], ensure_ascii=False) will_clauses.append(will_clause_json) # Optimal hops optimal_hops_list.append(scenario.get("optimal_hops", 3)) dataset = Dataset.from_dict({ "prompt": prompts, "gold_answer": gold_answers, "will_clause": will_clauses, "optimal_hops": optimal_hops_list, }) logger.info( "Dataset created: %d examples, columns=%s", len(dataset), dataset.column_names, ) # Persist --------------------------------------------------------------- if save_path: dataset.save_to_disk(save_path) logger.info("Dataset saved to %s", save_path) if push_to_hub: if not hub_repo_id: raise ValueError("hub_repo_id is required when push_to_hub=True") dataset.push_to_hub(hub_repo_id) logger.info("Dataset pushed to Hub: %s", hub_repo_id) return dataset # --------------------------------------------------------------------------- # Helpers # --------------------------------------------------------------------------- def _print_dataset_stats(dataset: Dataset) -> None: """Print summary statistics about a generated dataset.""" import statistics prompt_lengths = [len(p) for p in dataset["prompt"]] hop_values = dataset["optimal_hops"] # Will-type distribution will_types: dict[str, int] = {} for wc_json in dataset["will_clause"]: wc = json.loads(wc_json) clause = wc.get("clause_text", "") will_types[clause[:60] + "…"] = will_types.get(clause[:60] + "…", 0) + 1 print("=" * 60) print(" Sangue e Grafi — GRPO Dataset Statistics") print("=" * 60) print(f" Total examples: {len(dataset)}") print(f" Columns: {dataset.column_names}") print(f" Prompt length (chars): " f"min={min(prompt_lengths)}, " f"max={max(prompt_lengths)}, " f"mean={statistics.mean(prompt_lengths):.0f}") print(f" Optimal hops: " f"min={min(hop_values)}, " f"max={max(hop_values)}, " f"mean={statistics.mean(hop_values):.1f}") print(f" Unique gold answers: {len(set(dataset['gold_answer']))}") print() print(" Will clause distribution:") for clause_snippet, count in sorted( will_types.items(), key=lambda x: -x[1] ): print(f" {count:>4d} {clause_snippet}") print("=" * 60) # Show one example print() print(" Example prompt (first 500 chars):") print(" " + "-" * 56) example = dataset[0]["prompt"][:500] for line in example.split("\n"): print(f" | {line}") print(" " + "-" * 56) print(f" Gold answer: {dataset[0]['gold_answer']}") print() # --------------------------------------------------------------------------- # CLI entry-point # --------------------------------------------------------------------------- if __name__ == "__main__": logging.basicConfig( level=logging.INFO, format="%(asctime)s %(name)s %(levelname)s %(message)s", ) # Generate a small smoke-test dataset (using safe training seeds) ds = prepare_grpo_dataset( n_scenarios=50, seed=99990, # Smoke test range — NOT benchmark seeds save_path="data/grpo_smoke_test", ) _print_dataset_stats(ds)