Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
18 changes: 15 additions & 3 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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 `<your-router>` 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

Expand Down
2 changes: 2 additions & 0 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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
69 changes: 22 additions & 47 deletions router_inference/generate_prediction_file.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand All @@ -23,34 +23,15 @@
# 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",
"full": "./dataset/router_data.json",
}


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.
Expand All @@ -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 = {
Expand Down Expand Up @@ -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",
Expand All @@ -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...")
Expand All @@ -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...")
Expand Down
9 changes: 9 additions & 0 deletions router_inference/router/__init__.py
Original file line number Diff line number Diff line change
@@ -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"]
146 changes: 146 additions & 0 deletions router_inference/router/base_router.py
Original file line number Diff line number Diff line change
@@ -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
46 changes: 46 additions & 0 deletions router_inference/router/example_router.py
Original file line number Diff line number Diff line change
@@ -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