diff --git a/README.md b/README.md index c17b71c6..644ea1c3 100644 --- a/README.md +++ b/README.md @@ -128,15 +128,27 @@ For each model in your config, add an entry with the pricing per million tokens > [!NOTE] > Ensure all models in your above config files are listed in [`./universal_model_names.py`](./universal_model_names.py). If you add a new model, you must also add the API inference endpoint in [`llm_inference/model_inference.py`](./llm_inference/model_inference.py). -### Step 2.2: Generate Router's Prediction File +### Step 2.2: Create Your Router Class and Generate Prediction File -Generate a template prediction file: +Create your own router class by inheriting from `BaseRouter` and implementing the `_get_prediction()` method. See [`router_inference/router/example_router.py`](./router_inference/router/example_router.py) for a complete example. + +Then, modify [`router_inference/generate_prediction_file.py`](./router_inference/generate_prediction_file.py#L150) to use your router class: + +```python +# Replace ExampleRouter with your router class +from router_inference.router.my_router import MyRouter +router = MyRouter(args.router_name) +``` + +Finally, generate the prediction file: ```bash uv run python ./router_inference/generate_prediction_file.py your-router [sub_10|full] ``` -**Important**: Replace the placeholder model choices of the `prediction` field in the generated prediction file with your router's actual selections. We will automate this process in a future version. +> [!NOTE] +> - The `` argument must match your config filename (without the `.json` extension). For example, if your config file is `router_inference/config/my-router.json`, use `my-router` as the argument. +> - Your `_get_prediction()` method must return a model name that exists in your config file's `models` list. The base class will automatically validate this. ### Step 2.3: Validate Config and Prediction Files diff --git a/pyproject.toml b/pyproject.toml index 94be8e47..98144f83 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -106,3 +106,5 @@ plugins = ['pydantic.mypy'] ignore_missing_imports = true check_untyped_defs = true follow_imports = "silent" +namespace_packages = true +explicit_package_bases = true diff --git a/router_inference/generate_prediction_file.py b/router_inference/generate_prediction_file.py index e532cf02..5d9ef380 100644 --- a/router_inference/generate_prediction_file.py +++ b/router_inference/generate_prediction_file.py @@ -2,10 +2,10 @@ # SPDX-License-Identifier: Apache-2.0 """ -Generate Prediction File for Toy Router. +Generate Prediction File using ExampleRouter. -This script generates a prediction file for a toy router that cycles through -models in the config file using a simple modulo operation. This is useful for +This script generates a prediction file using the ExampleRouter class, +which cycles through models in the config file. This is useful for testing the RouterArena pipeline. Usage: @@ -23,6 +23,8 @@ # Add parent directory to path for imports sys.path.append(os.path.abspath(os.path.join(os.path.dirname(__file__), "../"))) +from router_inference.router import ExampleRouter, BaseRouter + # Dataset file paths DATASET_PATHS = { "sub_10": "./dataset/router_data_10.json", @@ -30,27 +32,6 @@ } -def load_config(router_name: str) -> Dict[str, Any]: - """ - Load router config file. - - Args: - router_name: Name of the router - - Returns: - Configuration dictionary - """ - config_path = f"./router_inference/config/{router_name}.json" - - if not os.path.exists(config_path): - raise FileNotFoundError(f"Config file not found: {config_path}") - - with open(config_path, "r", encoding="utf-8") as f: - config = json.load(f) - - return config - - def load_dataset(split: str) -> List[Dict[str, Any]]: """ Load dataset file. @@ -76,37 +57,29 @@ def load_dataset(split: str) -> List[Dict[str, Any]]: def generate_predictions( - dataset: List[Dict[str, Any]], models: List[str] + dataset: List[Dict[str, Any]], router: BaseRouter ) -> List[Dict[str, Any]]: """ - Generate predictions by cycling through models using modulo operation. + Generate predictions using the ExampleRouter. Args: dataset: List of dataset entries - models: List of model names to cycle through + router: ExampleRouter instance to use for predictions Returns: List of prediction dictionaries """ predictions = [] - # Get the list of model names (only those with non-null cost) - model_names = [model for model in models] - - if not model_names: - raise ValueError("No models with non-null cost found in config") - - for i, entry in enumerate(dataset): + for entry in dataset: global_index = entry.get("global index") prompt = entry.get("prompt_formatted") or entry.get("prompt") if not global_index or not prompt: continue - # Use modulo operation to cycle through models - # NOTE: This is a toy router, you could implement your own logic here, or create a prediction file in your own router loop. - model_index = i % len(model_names) - selected_model = model_names[model_index] + # Use the router to get prediction (validation is handled by BaseRouter) + selected_model = router.get_prediction(prompt) # Create prediction entry prediction_entry = { @@ -145,7 +118,7 @@ def save_predictions(predictions: List[Dict[str, Any]], router_name: str) -> Non def main(): """Main function to handle command line arguments and generate predictions.""" parser = argparse.ArgumentParser( - description="Generate prediction file for toy router" + description="Generate prediction file using ExampleRouter" ) parser.add_argument( "router_name", @@ -170,12 +143,14 @@ def main(): print(f"Dataset split: {args.split}") print("=" * 80) - # Load config - print("\n[1] Loading config...") - config = load_config(args.router_name) - models = config.get("pipeline_params", {}).get("models", []) - print(f"✓ Config loaded: {len(models)} models in config") - print(f" Models: {', '.join(models)}") + # Initialize router + print("\n[1] Initializing router...") + + ## You should replace ExampleRouter with your own router implementation. + router = ExampleRouter(args.router_name) + + print(f"✓ Router initialized: {router.router_name}") + print(f" Available models: {', '.join(router.models)}") # Load dataset print("\n[2] Loading dataset...") @@ -184,9 +159,9 @@ def main(): # Generate predictions print("\n[3] Generating predictions...") - predictions = generate_predictions(dataset, models) + predictions = generate_predictions(dataset, router) print(f"✓ Generated {len(predictions)} predictions") - print(" Using toy router logic: cycling through models") + print(" Using ExampleRouter: cycling through models") # Save predictions print("\n[4] Saving predictions...") diff --git a/router_inference/router/__init__.py b/router_inference/router/__init__.py new file mode 100644 index 00000000..22228fce --- /dev/null +++ b/router_inference/router/__init__.py @@ -0,0 +1,9 @@ +# SPDX-FileCopyrightText: Copyright contributors to the RouterArena project +# SPDX-License-Identifier: Apache-2.0 + +"""Router inference module for RouterArena.""" + +from router_inference.router.base_router import BaseRouter +from router_inference.router.example_router import ExampleRouter + +__all__ = ["BaseRouter", "ExampleRouter"] diff --git a/router_inference/router/base_router.py b/router_inference/router/base_router.py new file mode 100644 index 00000000..e30e000f --- /dev/null +++ b/router_inference/router/base_router.py @@ -0,0 +1,146 @@ +# SPDX-FileCopyrightText: Copyright contributors to the RouterArena project +# SPDX-License-Identifier: Apache-2.0 + +""" +Abstract base class for router implementations. + +All router implementations must inherit from this class and implement +the _get_prediction() method. The public get_prediction() method handles +validation automatically. +""" + +import json +import os +from abc import ABC, abstractmethod +from typing import Dict, Any, List + + +class BaseRouter(ABC): + """ + Abstract base class for router implementations. + + This class provides the foundation for all router implementations. + It handles config loading and validation, while requiring subclasses + to implement the core routing logic. + + Args: + router_name: Name of the router (used to load config file) + + Attributes: + router_name: Name of the router + config: Router configuration dictionary + models: List of available models from config + """ + + def __init__(self, router_name: str): + """ + Initialize the router with a router name. + + Args: + router_name: Name of the router (used to load config file) + + Raises: + FileNotFoundError: If config file doesn't exist + ValueError: If config file is invalid + """ + self.router_name = router_name + self.config = self._load_config() + self.models = self._extract_models() + + def _load_config(self) -> Dict[str, Any]: + """ + Load router configuration from JSON file. + + Returns: + Configuration dictionary + + Raises: + FileNotFoundError: If config file doesn't exist + ValueError: If config structure is invalid + """ + script_dir = os.path.dirname(os.path.abspath(__file__)) + project_root = os.path.dirname(os.path.dirname(script_dir)) + config_path = os.path.join( + project_root, "router_inference", "config", f"{self.router_name}.json" + ) + + if not os.path.exists(config_path): + raise FileNotFoundError(f"Config file not found: {config_path}") + + with open(config_path, "r", encoding="utf-8") as f: + config = json.load(f) + + # Validate config structure + if "pipeline_params" not in config: + raise ValueError( + f"Invalid config structure: missing 'pipeline_params' in {config_path}" + ) + + if "models" not in config["pipeline_params"]: + raise ValueError( + "Invalid config structure: missing 'models' in pipeline_params" + ) + + return config + + def _extract_models(self) -> List[str]: + """ + Extract list of models from config. + + Returns: + List of model names + """ + return self.config["pipeline_params"]["models"] + + def _validate_model(self, model_name: str) -> None: + """ + Validate that the selected model is in the config. + + Args: + model_name: Name of the model to validate + + Raises: + ValueError: If model is not in the config + """ + if model_name not in self.models: + raise ValueError( + f"Model '{model_name}' is not in the router config. " + f"Available models: {self.models}" + ) + + @abstractmethod + def _get_prediction(self, query: str) -> str: + """ + Get the model prediction for a given query (internal implementation). + + This is the core method that must be implemented by all router subclasses. + It should analyze the query and return the name of the model to use. + Subclasses should not validate the model - that is handled by get_prediction(). + + Args: + query: The input query string + + Returns: + Name of the selected model + """ + pass + + def get_prediction(self, query: str) -> str: + """ + Get the model prediction for a given query (public method with validation). + + This method calls the subclass's _get_prediction() implementation and + validates that the returned model is in the config. + + Args: + query: The input query string + + Returns: + Name of the selected model (validated to be in config) + + Raises: + ValueError: If the returned model is not in the config + """ + model_name = self._get_prediction(query) + self._validate_model(model_name) + return model_name diff --git a/router_inference/router/example_router.py b/router_inference/router/example_router.py new file mode 100644 index 00000000..54196579 --- /dev/null +++ b/router_inference/router/example_router.py @@ -0,0 +1,46 @@ +# SPDX-FileCopyrightText: Copyright contributors to the RouterArena project +# SPDX-License-Identifier: Apache-2.0 + +""" +Example router implementation. + +This is a simple example router that demonstrates how to implement +the BaseRouter abstract class. It selects the first model in the config +for all queries. +""" + +from router_inference.router.base_router import BaseRouter + + +class ExampleRouter(BaseRouter): + """ + Example router implementation. + + This router simply cycles through the models from the config for all queries. + This is intended as a demonstration and should be replaced with actual + routing logic in production implementations. + """ + + def __init__(self, router_name: str): + super().__init__(router_name) + self.counter = 0 + self.length = len(self.models) + + def _get_prediction(self, query: str) -> str: + """ + Get the model prediction for a given query (internal implementation). + + This example implementation cycles through models in the config. + In a real implementation, you would analyze the query and select + the most appropriate model. + + Args: + query: The input query string + + Returns: + Name of the selected model + """ + # Simple example: cycle through models + model_name = self.models[self.counter % self.length] + self.counter += 1 + return model_name