ministral3 (#14251)
Signed-off-by: Xinyuan Tong <xinyuantong.cs@gmail.com> Co-authored-by: Yueming Yuan <yy28@illinois.edu>
This commit is contained in:
co-authored by
Yueming Yuan
parent
c1006fd8a1
commit
6d37e70883
@@ -83,9 +83,9 @@ async def process_sample(
|
||||
assert image is not None
|
||||
image_path = sample["image_path"]
|
||||
extra_body = None if lora_path is None else {"lora_path": lora_path}
|
||||
response = await client.chat.completions.create(
|
||||
model="default",
|
||||
messages=[
|
||||
payload = {
|
||||
"model": "default",
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
@@ -95,11 +95,11 @@ async def process_sample(
|
||||
],
|
||||
}
|
||||
],
|
||||
temperature=0,
|
||||
max_completion_tokens=sampling_params["max_new_tokens"],
|
||||
max_tokens=sampling_params["max_new_tokens"],
|
||||
extra_body=extra_body,
|
||||
)
|
||||
"extra_body": extra_body,
|
||||
}
|
||||
if sampling_params:
|
||||
payload.update(sampling_params)
|
||||
response = await client.chat.completions.create(**payload)
|
||||
return sample, response.choices[0].message.content
|
||||
|
||||
|
||||
|
||||
@@ -36,7 +36,8 @@ class EvalArgs:
|
||||
profile: bool = False
|
||||
profile_number: int = 5
|
||||
concurrency: int = 1
|
||||
max_new_tokens: int = 30
|
||||
max_new_tokens: Optional[int] = None
|
||||
temperature: Optional[float] = None
|
||||
response_answer_regex: str = "(.*)"
|
||||
lora_path: Optional[str] = None
|
||||
|
||||
@@ -101,6 +102,12 @@ class EvalArgs:
|
||||
default=EvalArgs.max_new_tokens,
|
||||
help="Maximum number of new tokens to generate per sample.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--temperature",
|
||||
type=float,
|
||||
default=EvalArgs.temperature,
|
||||
help="Sampling temperature for generation.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--response-answer-regex",
|
||||
type=str,
|
||||
@@ -241,19 +248,21 @@ def prepare_samples(eval_args: EvalArgs):
|
||||
|
||||
|
||||
def get_sampling_params(eval_args):
|
||||
max_new_tokens = eval_args.max_new_tokens
|
||||
temperature = 0.001
|
||||
|
||||
extra_request_body = {}
|
||||
if eval_args.extra_request_body:
|
||||
extra_request_body = json.loads(eval_args.extra_request_body)
|
||||
|
||||
return {
|
||||
"temperature": temperature,
|
||||
"max_new_tokens": max_new_tokens,
|
||||
sampling_params = {
|
||||
**extra_request_body,
|
||||
}
|
||||
|
||||
if eval_args.max_new_tokens is not None and eval_args.max_new_tokens > 0:
|
||||
sampling_params.update({"max_completion_tokens": eval_args.max_new_tokens})
|
||||
|
||||
if eval_args.temperature is not None:
|
||||
sampling_params.update({"temperature": eval_args.temperature})
|
||||
|
||||
return sampling_params
|
||||
|
||||
|
||||
# ----------- Process Multi-choice -------------
|
||||
def parse_multi_choice_response(response, all_choices, index2ans):
|
||||
|
||||
Reference in New Issue
Block a user