| 36 | import aiohttp |
| 37 | |
| 38 | def parse_args(): |
| 39 | parser = argparse.ArgumentParser(description="Run Naive RAG for various datasets and models.") |
| 40 | parser.add_argument('--dataset_name', type=str, required=True, help="Name of the dataset to use.") |
| 41 | parser.add_argument('--split', type=str, required=True, help="Dataset split to use.") |
| 42 | parser.add_argument('--subset_num', type=int, default=-1, help="Number of examples to process. Defaults to all if not specified.") |
| 43 | parser.add_argument('--top_k', type=int, default=10, help="Number of top search results to retrieve.") |
| 44 | parser.add_argument('--max_doc_len', type=int, default=3000, help="Maximum length of each searched document.") |
| 45 | parser.add_argument('--model_name', type=str, default="QwQ-32B", help="Name of the model to use") |
| 46 | parser.add_argument('--api_base_url', type=str, required=True, help="Base URL for the API endpoint") |
| 47 | parser.add_argument('--aux_model_name', type=str, default="Qwen2.5-72B-Instruct", help="Name of the model to use") |
| 48 | parser.add_argument('--aux_api_base_url', type=str, required=True, help="Base URL for the API endpoint") |
| 49 | parser.add_argument('--use_jina', action='store_true', help="Whether to use Jina API for document fetching.") |
| 50 | parser.add_argument('--jina_api_key', type=str, default='None', help="Your Jina API Key to Fetch URL Content.") |
| 51 | parser.add_argument('--temperature', type=float, default=0.7, help="Sampling temperature.") |
| 52 | parser.add_argument('--top_p', type=float, default=0.8, help="Top-p sampling parameter.") |
| 53 | parser.add_argument('--top_k_sampling', type=int, default=20, help="Top-k sampling parameter.") |
| 54 | parser.add_argument('--repetition_penalty', type=float, default=None, help="Repetition penalty. If not set, defaults based on the model.") |
| 55 | parser.add_argument('--max_tokens', type=int, default=32768, help="Maximum number of tokens to generate. If not set, defaults based on the model and dataset.") |
| 56 | parser.add_argument('--bing_subscription_key', type=str, default=None, help="Bing Search API subscription key.") |
| 57 | parser.add_argument('--bing_endpoint', type=str, default="https://api.bing.microsoft.com/v7.0/search", help="Bing Search API endpoint.") |
| 58 | parser.add_argument('--serper_api_key', type=str, default=None, help="Google Serper API key.") |
| 59 | parser.add_argument('--search_engine', type=str, default="bing", choices=["bing", "serper"], help="Search engine to use (bing or serper). Default: bing") |
| 60 | parser.add_argument('--concurrent_limit', type=int, default=50, help="Maximum number of concurrent API calls") |
| 61 | parser.add_argument('--seed', type=int, default=42, help="Random seed for reproducibility") |
| 62 | parser.add_argument('--eval', action='store_true', help="Whether to run evaluation") |
| 63 | parser.add_argument('--apply_query_planning', action='store_true', help="Whether to apply query planning for search") |
| 64 | return parser.parse_args() |
| 65 | |
| 66 | async def generate_response( |
| 67 | client: AsyncOpenAI, |