perf: increase eval batch sizes and add max_seq_len
- humaneval/ifeval: default batch_size 64, add --max_seq_len=4096 - mmlu: batch 4 questions x 4 choices per forward, add --batch_size - ppl: default batch_size 64
This commit is contained in:
@@ -57,6 +57,7 @@ class EvalConfig:
|
||||
top_p: float = 0.95
|
||||
top_k: int = 50
|
||||
batch_size: int = 32
|
||||
max_seq_len: int = 4096
|
||||
test_timeout: float = 3.0
|
||||
test_workers: int = 8
|
||||
k_values: Tuple[int, ...] = (1, 10, 100)
|
||||
@@ -90,7 +91,9 @@ def save_json(path: str, data):
|
||||
json.dump(data, f, indent=2, ensure_ascii=False)
|
||||
|
||||
|
||||
def create_engine(param_path: str, batch_size: int) -> InferenceEngine:
|
||||
def create_engine(
|
||||
param_path: str, batch_size: int, max_seq_len: int
|
||||
) -> InferenceEngine:
|
||||
model = AutoModel.from_pretrained(param_path)
|
||||
tokenizer = AutoTokenizer.from_pretrained(param_path)
|
||||
model.to(device="cuda", dtype=torch.bfloat16)
|
||||
@@ -98,6 +101,7 @@ def create_engine(param_path: str, batch_size: int) -> InferenceEngine:
|
||||
model=model,
|
||||
tokenizer=tokenizer,
|
||||
max_batch_size=batch_size,
|
||||
max_seq_len=max_seq_len,
|
||||
)
|
||||
|
||||
|
||||
@@ -318,7 +322,7 @@ def run_pipeline(cfg: EvalConfig) -> Dict:
|
||||
if cfg.problem_indices:
|
||||
problems = [problems[i] for i in cfg.problem_indices if i < len(problems)]
|
||||
|
||||
engine = create_engine(cfg.param_path, cfg.batch_size)
|
||||
engine = create_engine(cfg.param_path, cfg.batch_size, cfg.max_seq_len)
|
||||
|
||||
try:
|
||||
generated = generate_all(engine, problems, cfg)
|
||||
@@ -357,7 +361,8 @@ def parse_args(argv: Optional[List[str]] = None) -> EvalConfig:
|
||||
p.add_argument("--temperature", type=float, default=0.8)
|
||||
p.add_argument("--top_p", type=float, default=0.95)
|
||||
p.add_argument("--top_k", type=int, default=50)
|
||||
p.add_argument("--batch_size", type=int, default=32)
|
||||
p.add_argument("--batch_size", type=int, default=64)
|
||||
p.add_argument("--max_seq_len", type=int, default=4096)
|
||||
p.add_argument("--test_workers", type=int, default=8)
|
||||
p.add_argument("--test_timeout", type=float, default=3.0)
|
||||
p.add_argument("--problems", type=int, nargs="+", default=None)
|
||||
@@ -375,6 +380,7 @@ def parse_args(argv: Optional[List[str]] = None) -> EvalConfig:
|
||||
top_p=args.top_p,
|
||||
top_k=args.top_k,
|
||||
batch_size=args.batch_size,
|
||||
max_seq_len=args.max_seq_len,
|
||||
test_workers=args.test_workers,
|
||||
test_timeout=args.test_timeout,
|
||||
problem_indices=args.problems,
|
||||
|
||||
@@ -509,7 +509,10 @@ def main():
|
||||
help="Number of samples per problem (best-of-n scoring)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--batch_size", type=int, default=1, help="Inference batch size"
|
||||
"--batch_size", type=int, default=64, help="Inference batch size"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--max_seq_len", type=int, default=4096, help="Max sequence length for KV cache"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--limit",
|
||||
@@ -542,6 +545,7 @@ def main():
|
||||
model=model,
|
||||
tokenizer=tokenizer,
|
||||
max_batch_size=args.batch_size,
|
||||
max_seq_len=args.max_seq_len,
|
||||
)
|
||||
|
||||
results = evaluate(
|
||||
|
||||
@@ -179,31 +179,53 @@ def apply_chat(
|
||||
)
|
||||
|
||||
|
||||
def choice_logprob(
|
||||
model, tokenizer, context_ids: list[int], choice_letter: str, device: str
|
||||
) -> float:
|
||||
choice_text = choice_letter
|
||||
choice_ids = tokenizer.encode(choice_text, add_special_tokens=False)
|
||||
input_ids = context_ids + choice_ids
|
||||
max_len = model.config.max_position_embeddings
|
||||
if len(input_ids) > max_len:
|
||||
overflow = len(input_ids) - max_len
|
||||
input_ids = input_ids[overflow:]
|
||||
ctx_len = len(input_ids) - len(choice_ids)
|
||||
else:
|
||||
ctx_len = len(context_ids)
|
||||
def choice_logprobs_batched(
|
||||
model,
|
||||
tokenizer,
|
||||
context_ids_list: list[list[int]],
|
||||
device: str,
|
||||
max_model_len: int,
|
||||
) -> list[dict[str, float]]:
|
||||
"""Compute log-probs for multiple questions x 4 choices in batches.
|
||||
|
||||
Returns a list of dicts: [{A: score, B: score, C: score, D: score}, ...]
|
||||
"""
|
||||
letters = ("A", "B", "C", "D")
|
||||
choice_ids_list = [tokenizer.encode(c, add_special_tokens=False) for c in letters]
|
||||
|
||||
all_inputs: list[tuple[int, int, list[int], int, list[int]]] = []
|
||||
for qi, context_ids in enumerate(context_ids_list):
|
||||
for ci, choice_ids in enumerate(choice_ids_list):
|
||||
input_ids = context_ids + choice_ids
|
||||
if len(input_ids) > max_model_len:
|
||||
overflow = len(input_ids) - max_model_len
|
||||
input_ids = input_ids[overflow:]
|
||||
ctx_len = len(input_ids) - len(choice_ids)
|
||||
else:
|
||||
ctx_len = len(context_ids)
|
||||
all_inputs.append((qi, ci, input_ids, ctx_len, choice_ids))
|
||||
|
||||
n = len(all_inputs)
|
||||
max_input_len = max(len(x[2]) for x in all_inputs)
|
||||
padded = torch.zeros(n, max_input_len, dtype=torch.long, device=device)
|
||||
mask = torch.zeros(n, max_input_len, dtype=torch.bool, device=device)
|
||||
for i, (_, _, ids, _, _) in enumerate(all_inputs):
|
||||
padded[i, : len(ids)] = torch.tensor(ids, dtype=torch.long, device=device)
|
||||
mask[i, : len(ids)] = True
|
||||
|
||||
input_tensor = torch.tensor([input_ids], device=device, dtype=torch.long)
|
||||
with torch.inference_mode():
|
||||
logits = model(input_tensor)["logits"][0]
|
||||
logits = model(padded, input_mask=mask)["logits"]
|
||||
|
||||
score = 0.0
|
||||
for i, tid in enumerate(choice_ids):
|
||||
pos = ctx_len - 1 + i
|
||||
if pos >= len(logits):
|
||||
break
|
||||
score += F.log_softmax(logits[pos], dim=-1)[tid].item()
|
||||
return score
|
||||
results = [{} for _ in range(len(context_ids_list))]
|
||||
for i, (qi, ci, _, ctx_len, choice_ids) in enumerate(all_inputs):
|
||||
score = 0.0
|
||||
for j, tid in enumerate(choice_ids):
|
||||
pos = ctx_len - 1 + j
|
||||
if pos >= logits.size(1):
|
||||
break
|
||||
score += F.log_softmax(logits[i, pos].float(), dim=-1)[tid].item()
|
||||
results[qi][letters[ci]] = score
|
||||
return results
|
||||
|
||||
|
||||
def _permute_choices(item: dict, rng: random.Random) -> tuple[dict, str]:
|
||||
@@ -233,25 +255,42 @@ def evaluate_subject(
|
||||
device: str,
|
||||
n_shot: int,
|
||||
seed: int = 0,
|
||||
batch_size: int = 16,
|
||||
) -> tuple[float, int, int]:
|
||||
rng = random.Random(seed) if seed >= 0 else None
|
||||
correct = 0
|
||||
total = 0
|
||||
for item in tqdm.tqdm(test_data, desc=f"{subject:40s}", leave=False):
|
||||
|
||||
context_ids_list = []
|
||||
answers = []
|
||||
for item in test_data:
|
||||
if rng is not None:
|
||||
permuted, answer = _permute_choices(item, rng)
|
||||
else:
|
||||
permuted, answer = item, item["answer"]
|
||||
raw_prompt = build_prompt(permuted["question"], permuted, subject)
|
||||
context = apply_chat(tokenizer, raw_prompt, n_shot, dev_data or [], subject)
|
||||
context_ids = tokenizer.encode(context)
|
||||
scores = {
|
||||
c: choice_logprob(model, tokenizer, context_ids, c, device)
|
||||
for c in ("A", "B", "C", "D")
|
||||
}
|
||||
if max(scores, key=scores.get) == answer:
|
||||
correct += 1
|
||||
total += 1
|
||||
context_ids_list.append(tokenizer.encode(context))
|
||||
answers.append(answer)
|
||||
|
||||
max_model_len = model.config.max_position_embeddings
|
||||
|
||||
num_batches = (len(context_ids_list) + batch_size - 1) // batch_size
|
||||
for start in tqdm.tqdm(
|
||||
range(0, len(context_ids_list), batch_size),
|
||||
total=num_batches,
|
||||
desc=f"{subject:40s}",
|
||||
leave=False,
|
||||
):
|
||||
batch = context_ids_list[start : start + batch_size]
|
||||
batch_answers = answers[start : start + batch_size]
|
||||
scores_list = choice_logprobs_batched(
|
||||
model, tokenizer, batch, device, max_model_len
|
||||
)
|
||||
for scores, answer in zip(scores_list, batch_answers):
|
||||
if max(scores, key=scores.get) == answer:
|
||||
correct += 1
|
||||
total += 1
|
||||
return correct / total, correct, total
|
||||
|
||||
|
||||
@@ -290,6 +329,12 @@ def main():
|
||||
default=0,
|
||||
help="Seed for option permutation (0 to enable, -1 to disable)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--batch_size",
|
||||
type=int,
|
||||
default=4,
|
||||
help="Number of questions per batch (4 choices each = 4*B rows)",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
if args.download or not os.path.exists(args.data_dir):
|
||||
@@ -329,6 +374,7 @@ def main():
|
||||
device,
|
||||
args.n_shot,
|
||||
seed=args.seed,
|
||||
batch_size=args.batch_size,
|
||||
)
|
||||
results[subject] = {"accuracy": round(acc, 4), "correct": corr, "total": tot}
|
||||
total_correct += corr
|
||||
|
||||
@@ -415,7 +415,7 @@ if __name__ == "__main__":
|
||||
help="Key for the text field in the input data.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--batch_size", type=int, default=4, help="Batch size for evaluation."
|
||||
"--batch_size", type=int, default=64, help="Batch size for evaluation."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--max_length",
|
||||
|
||||
Reference in New Issue
Block a user