swaleha19/agent_tuning_framework
0
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 