Add CompletionSampler for non-chat eval in run_eval (#21785)

This commit is contained in:
Liangsheng Yin
2026-03-31 16:33:07 -07:00
committed by GitHub
parent 1f7cee81da
commit 7581d814ae
2 changed files with 85 additions and 3 deletions
+19 -3
View File
@@ -10,6 +10,7 @@ import time
from sglang.test.simple_eval_common import (
ChatCompletionSampler,
CompletionSampler,
Eval,
make_report,
set_ulimit,
@@ -60,16 +61,24 @@ def run_eval_once(args, base_url: str, eval_obj: Eval) -> dict:
if value is not None:
extra_body[param_name] = value
sampler = ChatCompletionSampler(
common_kwargs = dict(
model=args.model,
max_tokens=getattr(args, "max_tokens", 2048),
top_p=getattr(args, "top_p", 1.0),
base_url=base_url,
temperature=getattr(args, "temperature", 0.0),
reasoning_effort=getattr(args, "reasoning_effort", None),
extra_body=extra_body if extra_body else None,
)
api_mode = getattr(args, "api", "chat")
if api_mode == "completion":
sampler = CompletionSampler(**common_kwargs)
else:
sampler = ChatCompletionSampler(
**common_kwargs,
reasoning_effort=getattr(args, "reasoning_effort", None),
extra_body=extra_body if extra_body else None,
)
# Run eval
tic = time.perf_counter()
result = eval_obj(sampler)
@@ -266,6 +275,13 @@ if __name__ == "__main__":
"--repeat", type=int, default=1, help="repeat the evaluation n times"
)
parser.add_argument("--eval-name", type=str, default="mmlu")
parser.add_argument(
"--api",
type=str,
default="chat",
choices=["chat", "completion"],
help="API mode: 'chat' for /v1/chat/completions, 'completion' for /v1/completions",
)
parser.add_argument("--num-examples", type=int)
parser.add_argument("--num-threads", type=int, default=512)
parser.add_argument("--max-tokens", type=int, default=2048)
+66
View File
@@ -169,6 +169,72 @@ class ChatCompletionSampler(SamplerBase):
return ""
class CompletionSampler(SamplerBase):
"""
Sample from OpenAI's completion API (non-chat).
Sends raw text prompts without chat template wrapping.
"""
def __init__(
self,
base_url: str = None,
model: Optional[str] = None,
temperature: float = 0.0,
top_p: float = 1.0,
max_tokens: int = 2048,
):
self.client = OpenAI(base_url=base_url, http_client=LargerHttpxClient())
if model is None:
model = self.client.models.list().data[0].id
self.model = model
self.temperature = temperature
self.top_p = top_p
self.max_tokens = max_tokens
self._completion_tokens: list[int] = []
print(
f"CompletionSampler initialized with {self.model=} {self.temperature=} {self.max_tokens=}"
)
def _pack_message(self, role: str, content: Any):
return {"role": str(role), "content": content}
def __call__(self, message_list: MessageList) -> str:
# Extract raw text from message list (eval objects pack prompt as a single user message)
prompt = "\n".join(
msg["content"]
for msg in message_list
if isinstance(msg.get("content"), str)
)
trial = 0
while trial < 6:
try:
response = self.client.completions.create(
model=self.model,
prompt=prompt,
temperature=self.temperature,
top_p=self.top_p,
max_tokens=self.max_tokens,
)
if response.usage and response.usage.completion_tokens is not None:
self._completion_tokens.append(response.usage.completion_tokens)
return response.choices[0].text or ""
except openai.BadRequestError as e:
print("Bad Request Error", e)
return ""
except Exception as e:
exception_backoff = 2**trial
print(
f"Rate limit exception so wait and retry {trial} after {exception_backoff} sec",
e,
)
time.sleep(exception_backoff)
trial += 1
print(f"All retry attempts exhausted for request. Returning empty response.")
return ""
QUERY_TEMPLATE_MULTICHOICE = """
Answer the following multiple choice question. The last line of your response should be of the following format: 'Answer: $LETTER' (without quotes) where LETTER is one of ABCD. Think step by step before answering.