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_p: float = 0.95
|
||||||
top_k: int = 50
|
top_k: int = 50
|
||||||
batch_size: int = 32
|
batch_size: int = 32
|
||||||
|
max_seq_len: int = 4096
|
||||||
test_timeout: float = 3.0
|
test_timeout: float = 3.0
|
||||||
test_workers: int = 8
|
test_workers: int = 8
|
||||||
k_values: Tuple[int, ...] = (1, 10, 100)
|
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)
|
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)
|
model = AutoModel.from_pretrained(param_path)
|
||||||
tokenizer = AutoTokenizer.from_pretrained(param_path)
|
tokenizer = AutoTokenizer.from_pretrained(param_path)
|
||||||
model.to(device="cuda", dtype=torch.bfloat16)
|
model.to(device="cuda", dtype=torch.bfloat16)
|
||||||
@@ -98,6 +101,7 @@ def create_engine(param_path: str, batch_size: int) -> InferenceEngine:
|
|||||||
model=model,
|
model=model,
|
||||||
tokenizer=tokenizer,
|
tokenizer=tokenizer,
|
||||||
max_batch_size=batch_size,
|
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:
|
if cfg.problem_indices:
|
||||||
problems = [problems[i] for i in cfg.problem_indices if i < len(problems)]
|
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:
|
try:
|
||||||
generated = generate_all(engine, problems, cfg)
|
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("--temperature", type=float, default=0.8)
|
||||||
p.add_argument("--top_p", type=float, default=0.95)
|
p.add_argument("--top_p", type=float, default=0.95)
|
||||||
p.add_argument("--top_k", type=int, default=50)
|
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_workers", type=int, default=8)
|
||||||
p.add_argument("--test_timeout", type=float, default=3.0)
|
p.add_argument("--test_timeout", type=float, default=3.0)
|
||||||
p.add_argument("--problems", type=int, nargs="+", default=None)
|
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_p=args.top_p,
|
||||||
top_k=args.top_k,
|
top_k=args.top_k,
|
||||||
batch_size=args.batch_size,
|
batch_size=args.batch_size,
|
||||||
|
max_seq_len=args.max_seq_len,
|
||||||
test_workers=args.test_workers,
|
test_workers=args.test_workers,
|
||||||
test_timeout=args.test_timeout,
|
test_timeout=args.test_timeout,
|
||||||
problem_indices=args.problems,
|
problem_indices=args.problems,
|
||||||
|
|||||||
@@ -509,7 +509,10 @@ def main():
|
|||||||
help="Number of samples per problem (best-of-n scoring)",
|
help="Number of samples per problem (best-of-n scoring)",
|
||||||
)
|
)
|
||||||
parser.add_argument(
|
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(
|
parser.add_argument(
|
||||||
"--limit",
|
"--limit",
|
||||||
@@ -542,6 +545,7 @@ def main():
|
|||||||
model=model,
|
model=model,
|
||||||
tokenizer=tokenizer,
|
tokenizer=tokenizer,
|
||||||
max_batch_size=args.batch_size,
|
max_batch_size=args.batch_size,
|
||||||
|
max_seq_len=args.max_seq_len,
|
||||||
)
|
)
|
||||||
|
|
||||||
results = evaluate(
|
results = evaluate(
|
||||||
|
|||||||
@@ -179,31 +179,53 @@ def apply_chat(
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def choice_logprob(
|
def choice_logprobs_batched(
|
||||||
model, tokenizer, context_ids: list[int], choice_letter: str, device: str
|
model,
|
||||||
) -> float:
|
tokenizer,
|
||||||
choice_text = choice_letter
|
context_ids_list: list[list[int]],
|
||||||
choice_ids = tokenizer.encode(choice_text, add_special_tokens=False)
|
device: str,
|
||||||
input_ids = context_ids + choice_ids
|
max_model_len: int,
|
||||||
max_len = model.config.max_position_embeddings
|
) -> list[dict[str, float]]:
|
||||||
if len(input_ids) > max_len:
|
"""Compute log-probs for multiple questions x 4 choices in batches.
|
||||||
overflow = len(input_ids) - max_len
|
|
||||||
input_ids = input_ids[overflow:]
|
Returns a list of dicts: [{A: score, B: score, C: score, D: score}, ...]
|
||||||
ctx_len = len(input_ids) - len(choice_ids)
|
"""
|
||||||
else:
|
letters = ("A", "B", "C", "D")
|
||||||
ctx_len = len(context_ids)
|
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():
|
with torch.inference_mode():
|
||||||
logits = model(input_tensor)["logits"][0]
|
logits = model(padded, input_mask=mask)["logits"]
|
||||||
|
|
||||||
score = 0.0
|
results = [{} for _ in range(len(context_ids_list))]
|
||||||
for i, tid in enumerate(choice_ids):
|
for i, (qi, ci, _, ctx_len, choice_ids) in enumerate(all_inputs):
|
||||||
pos = ctx_len - 1 + i
|
score = 0.0
|
||||||
if pos >= len(logits):
|
for j, tid in enumerate(choice_ids):
|
||||||
break
|
pos = ctx_len - 1 + j
|
||||||
score += F.log_softmax(logits[pos], dim=-1)[tid].item()
|
if pos >= logits.size(1):
|
||||||
return score
|
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]:
|
def _permute_choices(item: dict, rng: random.Random) -> tuple[dict, str]:
|
||||||
@@ -233,25 +255,42 @@ def evaluate_subject(
|
|||||||
device: str,
|
device: str,
|
||||||
n_shot: int,
|
n_shot: int,
|
||||||
seed: int = 0,
|
seed: int = 0,
|
||||||
|
batch_size: int = 16,
|
||||||
) -> tuple[float, int, int]:
|
) -> tuple[float, int, int]:
|
||||||
rng = random.Random(seed) if seed >= 0 else None
|
rng = random.Random(seed) if seed >= 0 else None
|
||||||
correct = 0
|
correct = 0
|
||||||
total = 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:
|
if rng is not None:
|
||||||
permuted, answer = _permute_choices(item, rng)
|
permuted, answer = _permute_choices(item, rng)
|
||||||
else:
|
else:
|
||||||
permuted, answer = item, item["answer"]
|
permuted, answer = item, item["answer"]
|
||||||
raw_prompt = build_prompt(permuted["question"], permuted, subject)
|
raw_prompt = build_prompt(permuted["question"], permuted, subject)
|
||||||
context = apply_chat(tokenizer, raw_prompt, n_shot, dev_data or [], subject)
|
context = apply_chat(tokenizer, raw_prompt, n_shot, dev_data or [], subject)
|
||||||
context_ids = tokenizer.encode(context)
|
context_ids_list.append(tokenizer.encode(context))
|
||||||
scores = {
|
answers.append(answer)
|
||||||
c: choice_logprob(model, tokenizer, context_ids, c, device)
|
|
||||||
for c in ("A", "B", "C", "D")
|
max_model_len = model.config.max_position_embeddings
|
||||||
}
|
|
||||||
if max(scores, key=scores.get) == answer:
|
num_batches = (len(context_ids_list) + batch_size - 1) // batch_size
|
||||||
correct += 1
|
for start in tqdm.tqdm(
|
||||||
total += 1
|
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
|
return correct / total, correct, total
|
||||||
|
|
||||||
|
|
||||||
@@ -290,6 +329,12 @@ def main():
|
|||||||
default=0,
|
default=0,
|
||||||
help="Seed for option permutation (0 to enable, -1 to disable)",
|
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()
|
args = parser.parse_args()
|
||||||
|
|
||||||
if args.download or not os.path.exists(args.data_dir):
|
if args.download or not os.path.exists(args.data_dir):
|
||||||
@@ -329,6 +374,7 @@ def main():
|
|||||||
device,
|
device,
|
||||||
args.n_shot,
|
args.n_shot,
|
||||||
seed=args.seed,
|
seed=args.seed,
|
||||||
|
batch_size=args.batch_size,
|
||||||
)
|
)
|
||||||
results[subject] = {"accuracy": round(acc, 4), "correct": corr, "total": tot}
|
results[subject] = {"accuracy": round(acc, 4), "correct": corr, "total": tot}
|
||||||
total_correct += corr
|
total_correct += corr
|
||||||
|
|||||||
@@ -415,7 +415,7 @@ if __name__ == "__main__":
|
|||||||
help="Key for the text field in the input data.",
|
help="Key for the text field in the input data.",
|
||||||
)
|
)
|
||||||
parser.add_argument(
|
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(
|
parser.add_argument(
|
||||||
"--max_length",
|
"--max_length",
|
||||||
|
|||||||
Reference in New Issue
Block a user