Commit 1e69d4ee authored by YaningGao's avatar YaningGao
Browse files

Merge branch 'main' of github.com:RAGEN-AI/vagen into prompt_dev

parents db97d179 49fb248f
Loading
Loading
Loading
Loading
+8 −0
Original line number Diff line number Diff line
models:
  # qwen_0.5b:
    # provider: vllm
    # model_name: Qwen/Qwen2.5-0.5B-Instruct
    # max_tokens: 1024
    # temperature: 0.7
    # tensor_parallel_size: 1
    # gpu_memory_utilization: 0.9
models:
  # qwen_0.5b:
    # provider: vllm
+12 −0
Original line number Diff line number Diff line
# vagen/inference/run_inference.py

import os
import sys
import argparse
import logging
import yaml
@@ -7,6 +9,9 @@ import wandb
import pandas as pd
import numpy as np
from datetime import datetime
import pandas as pd
import numpy as np
from datetime import datetime
from typing import Dict, List, Any
from collections import defaultdict

@@ -19,6 +24,7 @@ logger = logging.getLogger(__name__)
def parse_args():
    """Parse command line arguments."""
    parser = argparse.ArgumentParser(description="Run inference with models")
    parser = argparse.ArgumentParser(description="Run inference with models")
    
    parser.add_argument("--inference_config_path", type=str, required=True,
                       help="Path to inference configuration YAML")
@@ -120,6 +126,7 @@ def log_results_to_wandb(results: List[Dict], inference_config: Dict) -> None:
    wandb.log(summary_metrics)

def main():
    """Main entry point for inference."""
    """Main entry point for inference."""
    args = parse_args()
    
@@ -133,12 +140,16 @@ def main():
        format='%(asctime)s - %(name)s - %(levelname)s - %(message)s'
    )
    
    logger.info("Starting inference pipeline")
    logger.info("Starting inference pipeline")
    
    # Load environment configurations
    env_configs = load_environment_configs_from_parquet(args.val_files_path)
    # Load environment configurations
    env_configs = load_environment_configs_from_parquet(args.val_files_path)
    logger.info(f"Loaded {len(env_configs)} environment configurations")
    
    # Process each model
    # Process each model
    models = model_config.get('models', {})
    for model_name, model_cfg in models.items():
@@ -168,6 +179,7 @@ def main():
            service.run(max_steps=inference_config.get('max_steps', 10))
            results = service.recording_to_log()
            
            # Log results to wandb
            # Log results to wandb
            if inference_config.get('use_wandb', True):
                log_results_to_wandb(results, inference_config)
+1 −1

File changed.

Contains only whitespace changes.