CoolFace
Apppublic

swaleha19/agent_tuning_framework

sourceHugging Facemitupdated 1y agoView on Hugging Face
0likes
main.py266 linesDownload Raw Back to root
1"""2Main Integration Module for Agent Tuning Optimization Framework3 4This module provides functionality for integrating all components of the framework5and running end-to-end experiments.6"""7 8import os9import json10import argparse11from typing import List, Dict, Any, Union, Optional, Tuple12 13from models.llm_interface import LLMInterface14from data.trajectory_data import Trajectory, TrajectoryDataset, create_synthetic_dataset15from training.negative_samples import create_negative_sample_generator16from training.synthetic_trajectories import create_synthetic_trajectory_generator17from training.agent_tuner import create_agent_tuner18from evaluation.evaluators import create_agent_evaluator19 20def run_experiment(21    experiment_config: Dict[str, Any],22    output_dir: str23) -> Dict[str, Any]:24    """25    Run an end-to-end experiment with the framework.26    27    Args:28        experiment_config: Experiment configuration29        output_dir: Directory to save results30        31    Returns:32        Dictionary of experiment results33    """34    print(f"Starting experiment: {experiment_config['name']}")35    36    # Create output directory37    os.makedirs(output_dir, exist_ok=True)38    39    # Save experiment configuration40    with open(f"{output_dir}/experiment_config.json", "w") as f:41        json.dump(experiment_config, f, indent=2)42    43    # Initialize LLM interface44    print("Initializing LLM interface...")45    llm_config = experiment_config.get("llm", {})46    llm_interface = LLMInterface(47        model_name=llm_config.get("model_name", "gpt2"),48        model_type=llm_config.get("model_type", "causal"),49        device=llm_config.get("device", "cpu"),50        max_length=llm_config.get("max_length", 512),51        temperature=llm_config.get("temperature", 0.7)52    )53    54    # Load or create dataset55    print("Preparing dataset...")56    dataset_config = experiment_config.get("dataset", {})57    58    if dataset_config.get("path"):59        # Load existing dataset60        dataset = TrajectoryDataset(dataset_config.get("name", "experiment_dataset"))61        dataset.load_from_json(dataset_config["path"])62    else:63        # Create synthetic dataset64        dataset = create_synthetic_dataset(dataset_config.get("num_trajectories", 20))65    66    print(f"Dataset loaded with {len(dataset.trajectories)} trajectories")67    68    # Generate negative samples69    print("Generating negative samples...")70    negative_config = experiment_config.get("negative_samples", {})71    72    if negative_config.get("enabled", True):73        negative_generator = create_negative_sample_generator(74            negative_config.get("method", "response_degradation")75        )76        77        positive_trajectories = dataset.get_trajectories(positive_only=True)78        negative_trajectories = negative_generator.batch_generate(79            positive_trajectories,80            **negative_config.get("params", {})81        )82        83        # Add negative trajectories to dataset84        for trajectory in negative_trajectories:85            dataset.add_trajectory(trajectory)86        87        print(f"Added {len(negative_trajectories)} negative trajectories")88    89    # Generate synthetic trajectories90    print("Generating synthetic trajectories...")91    synthetic_config = experiment_config.get("synthetic_trajectories", {})92    93    if synthetic_config.get("enabled", True):94        synthetic_generator = create_synthetic_trajectory_generator(95            synthetic_config.get("method", "template"),96            llm_interface if synthetic_config.get("method") in ["llm", "hybrid"] else None97        )98        99        # Generate from task descriptions100        task_descriptions = [t.task_description for t in dataset.get_trajectories(positive_only=True)]101        task_descriptions = list(set(task_descriptions))  # Remove duplicates102        103        synthetic_trajectories = synthetic_generator.batch_generate(104            task_descriptions,105            **synthetic_config.get("params", {})106        )107        108        # Add synthetic trajectories to dataset109        for trajectory in synthetic_trajectories:110            dataset.add_trajectory(trajectory)111        112        print(f"Added {len(synthetic_trajectories)} synthetic trajectories")113    114    # Save the enhanced dataset115    dataset.save_to_json(f"{output_dir}/enhanced_dataset.json")116    117    # Analyze dataset118    dataset_stats = dataset.analyze_dataset()119    with open(f"{output_dir}/dataset_stats.json", "w") as f:120        json.dump(dataset_stats, f, indent=2)121    122    # Split dataset for training and evaluation123    all_trajectories = dataset.get_trajectories()124    split_idx = int(len(all_trajectories) * 0.8)  # 80% for training125    126    train_trajectories = all_trajectories[:split_idx]127    eval_trajectories = all_trajectories[split_idx:]128    129    print(f"Split dataset: {len(train_trajectories)} for training, {len(eval_trajectories)} for evaluation")130    131    # Tune agent132    print("Tuning agent...")133    tuning_config = experiment_config.get("tuning", {})134    135    tuner = create_agent_tuner(tuning_config.get("method", "supervised"))136    137    tuned_model, tuning_metrics = tuner.tune(138        model_name=llm_config.get("model_name", "gpt2"),139        trajectories=train_trajectories,140        output_dir=f"{output_dir}/tuned_model",141        **tuning_config.get("params", {})142    )143    144    # Save tuning metrics145    with open(f"{output_dir}/tuning_metrics.json", "w") as f:146        # Convert any non-serializable values to strings147        serializable_metrics = {}148        for k, v in tuning_metrics.items():149            if isinstance(v, (int, float, str, bool, list, dict)) or v is None:150                serializable_metrics[k] = v151            else:152                serializable_metrics[k] = str(v)153        154        json.dump(serializable_metrics, f, indent=2)155    156    # Create tuned model interface157    tuned_llm_interface = LLMInterface(158        model_name=f"{output_dir}/tuned_model",159        model_type=llm_config.get("model_type", "causal"),160        device=llm_config.get("device", "cpu"),161        max_length=llm_config.get("max_length", 512),162        temperature=llm_config.get("temperature", 0.7)163    )164    165    # Evaluate agent166    print("Evaluating agent...")167    eval_config = experiment_config.get("evaluation", {})168    169    evaluator = create_agent_evaluator(eval_config.get("method", "quality"))170    171    eval_results = evaluator.evaluate(172        llm_interface=tuned_llm_interface,173        test_trajectories=eval_trajectories,174        **eval_config.get("params", {})175    )176    177    # Visualize evaluation results178    evaluator.visualize_results(179        results=eval_results,180        output_dir=f"{output_dir}/evaluation"181    )182    183    # Save evaluation results184    with open(f"{output_dir}/evaluation_results.json", "w") as f:185        # Create a simplified version without large data186        simplified_results = {}187        188        if "aggregated" in eval_results:189            simplified_results["aggregated"] = eval_results["aggregated"]190        191        if "metrics" in eval_results:192            # Include only essential metrics193            simplified_results["metrics"] = [194                {k: v for k, v in m.items() if k not in ["generated_responses"]}195                for m in eval_results["metrics"]196            ]197        198        json.dump(simplified_results, f, indent=2)199    200    # Comparative evaluation (if configured)201    if eval_config.get("comparative", {}).get("enabled", False):202        print("Performing comparative evaluation...")203        204        # Create baseline model interface205        baseline_llm_interface = LLMInterface(206            model_name=llm_config.get("model_name", "gpt2"),207            model_type=llm_config.get("model_type", "causal"),208            device=llm_config.get("device", "cpu"),209            max_length=llm_config.get("max_length", 512),210            temperature=llm_config.get("temperature", 0.7)211        )212        213        # Create comparative evaluator214        comparative_evaluator = create_agent_evaluator("comparative")215        216        # Evaluate and compare217        comparative_results = comparative_evaluator.evaluate(218            llm_interfaces={219                "baseline": baseline_llm_interface,220                "tuned": tuned_llm_interface221            },222            test_trajectories=eval_trajectories,223            **eval_config.get("comparative", {}).get("params", {})224        )225        226        # Visualize comparative results227        comparative_evaluator.visualize_results(228            results=comparative_results,229            output_dir=f"{output_dir}/comparative"230        )231        232        # Save comparative results233        with open(f"{output_dir}/comparative_results.json", "w") as f:234            # Create a simplified version235            simplified_comparative = {236                "comparative": comparative_results.get("comparative", {})237            }238            239            json.dump(simplified_comparative, f, indent=2)240    241    print(f"Experiment completed. Results saved to {output_dir}")242    243    return {244        "dataset_stats": dataset_stats,245        "tuning_metrics": tuning_metrics,246        "evaluation_results": eval_results247    }248 249def main():250    """Main function for running the framework from command line."""251    parser = argparse.ArgumentParser(description="Agent Tuning Optimization Framework")252    parser.add_argument("--config", type=str, required=True, help="Path to experiment configuration file")253    parser.add_argument("--output", type=str, default="./experiment_results", help="Directory to save results")254    255    args = parser.parse_args()256    257    # Load experiment configuration258    with open(args.config, "r") as f:259        experiment_config = json.load(f)260    261    # Run experiment262    run_experiment(experiment_config, args.output)263 264if __name__ == "__main__":265    main()266