diff --git a/README.md b/README.md index 0b4e2b3..a0bf52f 100644 --- a/README.md +++ b/README.md @@ -49,9 +49,16 @@ pip install -e . python scripts/generate_kernel_and_verify.py \ --op-name aten::add \ --single-test \ - --server-type openai + --api-format openai \ + --model-name your-model \ + --base-url https://your-provider.example/v1 \ + --api-key your-key ``` +`--api-format` describes the API protocol, not a registered provider. Any +OpenAI-compatible or Anthropic-compatible endpoint can be used directly with +`--base-url`, `--api-key`, and `--model-name`. + 👉 **For detailed setup, see [Getting Started](docs/source/getting-started/index.md).** ## Documentation @@ -89,4 +96,4 @@ python scripts/generate_kernel_and_verify.py \ ## License -Apache 2.0 License \ No newline at end of file +Apache 2.0 License diff --git a/README.zh-CN.md b/README.zh-CN.md index f0df693..f149520 100644 --- a/README.zh-CN.md +++ b/README.zh-CN.md @@ -148,24 +148,41 @@ export OPENAI_BASE_URL=http://your-endpoint/v1 # 可选,自定义端点 python scripts/generate_kernel_and_verify.py \ --op-name aten::add \ --single-test \ - --server-type openai \ + --api-format openai \ --model-name your-model-name \ --max-rounds 3 # 完整测试(全部 210 个算子) python scripts/generate_kernel_and_verify.py \ - --server-type openai \ + --api-format openai \ --model-name your-model-name \ --max-rounds 3 # 非 NVIDIA 芯片(仅 ATen) python scripts/generate_kernel_and_verify.py \ --dataset KernelGenBench-aten \ - --server-type openai \ + --api-format openai \ --model-name your-model-name \ --max-rounds 3 ``` +`--api-format` 表示接口协议,而不是需要预先注册的服务商名称。任何 +OpenAI-compatible 或 Anthropic-compatible 接口都可以直接通过 +`--base-url`、`--api-key` 和 `--model-name` 使用。例如: + +```bash +python scripts/generate_kernel_and_verify.py \ + --op-name aten::add \ + --api-format anthropic \ + --model-name your-model-name \ + --base-url https://your-provider.example \ + --api-key your-key +``` + +也可以省略 `--base-url` 和 `--api-key`,分别使用 +`OPENAI_BASE_URL` / `OPENAI_API_KEY` 或 +`ANTHROPIC_BASE_URL` / `ANTHROPIC_API_KEY` 环境变量。 + ### 参数说明 | 参数 | 说明 | 默认值 | @@ -173,7 +190,7 @@ python scripts/generate_kernel_and_verify.py \ | `--op-name` | 指定单个算子(如 `aten::add`、`vllm13::rms_norm`) | 全部算子 | | `--single-test` | 随机选 1 个算子快速测试 | 关闭 | | `--dataset` | 数据集(`KernelGenBench`、`KernelGenBench-aten`、`-vllm`、`-cublas`) | 自动检测 | -| `--server-type` | LLM 提供商(`openai`、`anthropic`) | `openai` | +| `--api-format` | API 协议(`openai`、`anthropic`);`--server-type` 是兼容别名 | `openai` | | `--model-name` | 模型名称 | `gpt-4o` | | `--max-rounds` | Pass@K 轮数 | 10 | | `--device-count` | 验证使用的 GPU 数量 | 8 | diff --git a/docs/source/operation-guide/llm-track/commands.md b/docs/source/operation-guide/llm-track/commands.md index 8b53f72..e673fbe 100644 --- a/docs/source/operation-guide/llm-track/commands.md +++ b/docs/source/operation-guide/llm-track/commands.md @@ -74,13 +74,13 @@ python scripts/generate_kernel_and_verify.py \ --server-type openai ``` -## Server Types +## API Formats ### OpenAI ```bash python scripts/generate_kernel_and_verify.py \ - --server-type openai \ + --api-format openai \ --model-name gpt-4o ``` @@ -88,22 +88,26 @@ python scripts/generate_kernel_and_verify.py \ ```bash python scripts/generate_kernel_and_verify.py \ - --server-type anthropic \ + --api-format anthropic \ --model-name claude-opus-4-6 ``` -### Third-Party Providers +### Compatible Endpoints -Use `--base-url` to connect to any OpenAI-compatible provider. +No provider registration is required. Select the endpoint's wire protocol and +pass its URL, key, and model directly: ```bash python scripts/generate_kernel_and_verify.py \ - --server-type openai \ + --api-format \ --model-name \ --base-url \ --api-key ``` +The same values can be provided with `OPENAI_BASE_URL` / +`OPENAI_API_KEY` or `ANTHROPIC_BASE_URL` / `ANTHROPIC_API_KEY`. + ## Advanced Options ### Enable Reflection diff --git a/docs/source/operation-guide/llm-track/parameters.md b/docs/source/operation-guide/llm-track/parameters.md index d38de0f..7b94d04 100644 --- a/docs/source/operation-guide/llm-track/parameters.md +++ b/docs/source/operation-guide/llm-track/parameters.md @@ -22,7 +22,7 @@ LLM Track command-line parameters. | Parameter | Description | |-----------|-------------| -| `--server-type` | LLM provider: `openai` or `anthropic` | +| `--api-format` | API protocol: `openai` or `anthropic` | | `--model-name` | Model identifier | ## Optional Parameters @@ -31,7 +31,7 @@ LLM Track command-line parameters. |-----------|---------|-------------| | `--op-name` | All | Test a single operator (e.g., `aten::add`) | | `--single-test` | Off | Randomly select 1 operator for quick testing | -| `--base-url` | `http://localhost:8000/v1` | API base URL for OpenAI-compatible providers (e.g., DashScope, vLLM server) | +| `--base-url` | SDK default / Env var | API base URL for either compatible protocol | | `--api-key` | Env var | API key (overrides `OPENAI_API_KEY` / `ANTHROPIC_API_KEY` env var) | | `--dataset` | Auto | Dataset: `KernelGenBench`, `KernelGenBench-aten`, `KernelGenBench-vllm`, `KernelGenBench-cublas` | | `--max-rounds` | 10 | Number of Pass@K rounds | @@ -83,10 +83,10 @@ Number of independent kernel samples to generate: ### --base-url -Specify a custom API endpoint for OpenAI-compatible providers: +Specify a custom OpenAI-compatible or Anthropic-compatible endpoint: ```bash ---server-type openai --model-name --base-url +--api-format --model-name --base-url ``` ### --api-key @@ -97,7 +97,10 @@ Override the default API key from environment variables: --api-key ``` -If not set, reads from `OPENAI_API_KEY` or `ANTHROPIC_API_KEY` depending on `--server-type`. +If not set, the selected protocol reads `OPENAI_API_KEY` or +`ANTHROPIC_API_KEY`. The corresponding base URL can be supplied through +`OPENAI_BASE_URL` or `ANTHROPIC_BASE_URL`. `--server-type` remains +available as a backward-compatible alias for `--api-format`. ## Output diff --git a/scripts/generate_kernel_and_verify.py b/scripts/generate_kernel_and_verify.py index 49daea7..e601810 100644 --- a/scripts/generate_kernel_and_verify.py +++ b/scripts/generate_kernel_and_verify.py @@ -775,9 +775,15 @@ def main(): parser.add_argument("--timeout", type=int, default=300, help="Timeout for each test") # Generation config - parser.add_argument("--server-type", type=str, default="openai") + parser.add_argument( + "--api-format", "--server-type", + dest="server_type", + choices=["openai", "anthropic"], + default="openai", + help="API wire format; --server-type is kept as a backward-compatible alias", + ) parser.add_argument("--model-name", type=str, default="gpt-4o-mini") - parser.add_argument("--base-url", type=str, default=None, help="API base URL (for OpenAI-compatible providers)") + parser.add_argument("--base-url", type=str, default=None, help="API base URL for either supported API format") parser.add_argument("--api-key", type=str, default=None, help="API key (overrides OPENAI_API_KEY / ANTHROPIC_API_KEY env var)") parser.add_argument("--temperature", type=float, default=0.8) parser.add_argument("--max-tokens", type=int, default=16384) @@ -823,6 +829,8 @@ def main(): args_file = output_dir / "args.json" with open(args_file, "w") as f: args_dict = vars(args).copy() + if args_dict.get("api_key"): + args_dict["api_key"] = "" # Convert Path objects to strings for JSON serialization for key, value in args_dict.items(): if isinstance(value, Path): @@ -832,18 +840,18 @@ def main(): run_name = output_dir.name # Create generation config - # Set API key in env if provided + # Set only the environment variable for the selected wire format. The + # inference client reads it when each request is created, so CLI values work + # even though this module imports the generator before parsing arguments. if args.api_key: - os.environ["OPENAI_API_KEY"] = args.api_key - os.environ["ANTHROPIC_API_KEY"] = args.api_key - - base_url = args.base_url if args.base_url else "http://localhost:8000/v1" + key_env = "ANTHROPIC_API_KEY" if args.server_type == "anthropic" else "OPENAI_API_KEY" + os.environ[key_env] = args.api_key gen_config = GenerationConfig( run_name="", server_type=args.server_type, model_name=args.model_name, - base_url=base_url, + base_url=args.base_url, temperature=args.temperature, max_tokens=args.max_tokens, num_workers=args.num_workers, diff --git a/src/generator/sampler/generate_samples.py b/src/generator/sampler/generate_samples.py index 57871de..6b47cc1 100644 --- a/src/generator/sampler/generate_samples.py +++ b/src/generator/sampler/generate_samples.py @@ -82,7 +82,7 @@ class GenerationConfig: log_prompt: bool = False backend: str = "triton" greedy_sample: bool = False - base_url: str = "http://localhost:8000/v1" + base_url: Optional[str] = None strict_check: bool = False seed: int = 42 use_ai_advice: bool = False diff --git a/src/generator/sampler/utils.py b/src/generator/sampler/utils.py index aa805d4..b0ee359 100644 --- a/src/generator/sampler/utils.py +++ b/src/generator/sampler/utils.py @@ -24,10 +24,6 @@ logger = logging.getLogger(__name__) -ANTHROPIC_KEY = os.environ.get("ANTHROPIC_API_KEY") or os.environ.get("ANTHROPIC_AUTH_TOKEN") -ANTHROPIC_BASE_URL = os.environ.get("ANTHROPIC_BASE_URL") -OPENAI_KEY = os.environ.get("OPENAI_API_KEY") - ############################################ # Triton Prompt ############################################ @@ -86,20 +82,34 @@ def query_server( base_url: str = None, **kwargs, ): + if server_type not in {"anthropic", "openai"}: + raise ValueError( + f"Unsupported API format: {server_type!r}. " + "Use 'openai' for OpenAI-compatible APIs or 'anthropic' for " + "Anthropic-compatible APIs." + ) + match server_type: case "anthropic": import anthropic as _anthropic - client = _anthropic.Anthropic( - api_key=ANTHROPIC_KEY, - base_url=ANTHROPIC_BASE_URL if ANTHROPIC_BASE_URL else _anthropic.NOT_GIVEN, - ) + client_args = {} + api_key = os.environ.get("ANTHROPIC_API_KEY") or os.environ.get("ANTHROPIC_AUTH_TOKEN") + resolved_base_url = base_url or os.environ.get("ANTHROPIC_BASE_URL") + if api_key: + client_args["api_key"] = api_key + if resolved_base_url: + client_args["base_url"] = resolved_base_url + client = _anthropic.Anthropic(**client_args) model = model_name case "openai": - client = OpenAI(api_key=OPENAI_KEY) - model = model_name - case _: - _base_url = base_url or os.environ.get("OPENAI_BASE_URL", "http://localhost:8000/v1") - client = OpenAI(api_key=os.environ.get("OPENAI_API_KEY", "EMPTY"), base_url=_base_url) + client_args = {} + api_key = os.environ.get("OPENAI_API_KEY") + resolved_base_url = base_url or os.environ.get("OPENAI_BASE_URL") + if api_key: + client_args["api_key"] = api_key + if resolved_base_url: + client_args["base_url"] = resolved_base_url + client = OpenAI(**client_args) model = model_name if server_type == "anthropic": @@ -124,37 +134,27 @@ def query_server( max_tokens=max_tokens, ) outputs = [choice.text for choice in response.content if not hasattr(choice, 'thinking') or not choice.thinking] - elif server_type == "openai" and is_reasoning_model: - response = client.chat.completions.create( - model=model, - messages=[{"role": "user", "content": prompt}], - reasoning_effort=reasoning_effort, - ) - outputs = [choice.message.content for choice in response.choices] else: - if type(prompt) == str: - response = client.completions.create( + messages = prompt if isinstance(prompt, list) else [ + {"role": "system", "content": system_prompt}, + {"role": "user", "content": prompt}, + ] + if is_reasoning_model: + response = client.chat.completions.create( model=model, - prompt=prompt, - temperature=temperature, - n=num_completions, - max_tokens=max_tokens, - top_p=top_p, + messages=messages, + reasoning_effort=reasoning_effort, ) - outputs = [choice.text for choice in response.choices] else: response = client.chat.completions.create( model=model, - messages=prompt if isinstance(prompt, list) else [ - {"role": "system", "content": system_prompt}, - {"role": "user", "content": prompt}, - ], + messages=messages, temperature=temperature, n=num_completions, max_tokens=max_tokens, top_p=top_p, ) - outputs = [choice.message.content for choice in response.choices] + outputs = [choice.message.content for choice in response.choices] return outputs[0] if len(outputs) == 1 else outputs diff --git a/tests/test_sampler_api_config.py b/tests/test_sampler_api_config.py new file mode 100644 index 0000000..8d33b20 --- /dev/null +++ b/tests/test_sampler_api_config.py @@ -0,0 +1,69 @@ +from types import SimpleNamespace +from unittest.mock import MagicMock + +import pytest + +from generator.sampler import utils + + +def test_openai_compatible_endpoint_uses_runtime_env_and_chat(monkeypatch): + monkeypatch.setenv("OPENAI_API_KEY", "runtime-key") + monkeypatch.setenv("OPENAI_BASE_URL", "https://env.example/v1") + + client = MagicMock() + client.chat.completions.create.return_value = SimpleNamespace( + choices=[SimpleNamespace(message=SimpleNamespace(content="kernel"))] + ) + client_factory = MagicMock(return_value=client) + monkeypatch.setattr(utils, "OpenAI", client_factory) + + result = utils.query_server( + "generate a kernel", + server_type="openai", + model_name="custom-model", + base_url="https://cli.example/v1", + ) + + assert result == "kernel" + client_factory.assert_called_once_with( + api_key="runtime-key", + base_url="https://cli.example/v1", + ) + request = client.chat.completions.create.call_args.kwargs + assert request["model"] == "custom-model" + assert request["messages"] == [ + {"role": "system", "content": "You are a helpful assistant"}, + {"role": "user", "content": "generate a kernel"}, + ] + + +def test_anthropic_compatible_endpoint_uses_runtime_env(monkeypatch): + import anthropic + + monkeypatch.setenv("ANTHROPIC_API_KEY", "runtime-key") + monkeypatch.setenv("ANTHROPIC_BASE_URL", "https://env.example") + + client = MagicMock() + client.messages.create.return_value = SimpleNamespace( + content=[SimpleNamespace(text="kernel")] + ) + client_factory = MagicMock(return_value=client) + monkeypatch.setattr(anthropic, "Anthropic", client_factory) + + result = utils.query_server( + "generate a kernel", + server_type="anthropic", + model_name="custom-model", + base_url="https://cli.example", + ) + + assert result == "kernel" + client_factory.assert_called_once_with( + api_key="runtime-key", + base_url="https://cli.example", + ) + + +def test_provider_name_is_not_an_api_format(): + with pytest.raises(ValueError, match="OpenAI-compatible"): + utils.query_server("prompt", server_type="some-provider")