import os
import random
import re
import torch
import argparse
from vllm import LLM, SamplingParams
from huggingface_hub import snapshot_download
# --- GPU Helpers ---
def get_num_gpus():
"""Returns the number of available GPUs."""
return torch.cuda.device_count()
# --- Problem Generation & Parsing ---
def generate_multiplication_problems(n=1000):
"""Creates n random multiplication problems and their solutions."""
problems = []
solutions = []
for _ in range(n):
a = random.randint(2, 100)
b = random.randint(2, 100)
problems.append(f"What is {a} multiplied by {b}? Return only the solution.")
solutions.append(a * b)
return problems, solutions
def extract_last_number(text):
"""Extracts the last integer from a string (used to parse model outputs)."""
numbers = re.findall(r'-?\d+', text)
return int(numbers[-1]) if numbers else None
# --- Core Evaluation ---
def evaluate_multiplication(model_id, to_stdout=False):
"""
Evaluates how well a language model can solve multiplication problems.
Downloads the model, generates problems, queries the model, and calculates accuracy.
"""
# Download the model locally from Hugging Face
model_path = snapshot_download(model_id)
# Decide tensor parallelism based on available GPUs
num_gpus = get_num_gpus()
tensor_parallel_size = max(1, num_gpus) # Use TP only if more than 1 GPU is available
# Initialize the vLLM engine
llm = LLM(
model=model_path,
tensor_parallel_size=tensor_parallel_size
)
# Sampling configuration for deterministic outputs
sampling_params = SamplingParams(temperature=0.0, max_tokens=768)
# Generate multiplication problems and their correct answers
problems, solutions = generate_multiplication_problems()
# Query the model for answers
outputs = llm.generate(problems, sampling_params)
# Extract numeric predictions from model responses
predictions = []
for out in outputs:
text = out.outputs[0].text
predicted = extract_last_number(text)
predictions.append(predicted)
# Compute accuracy
correct = sum(1 for pred, sol in zip(predictions, solutions) if pred == sol)
accuracy = correct / len(solutions)
if to_stdout:
print(f"Model {model_id} accuracy: {accuracy:.4f}")
else:
os.makedirs('accuracies', exist_ok=True)
with open(f'accuracies/{model_id.replace("/", "_")}.txt', 'w') as f:
f.write(f"{accuracy:.4f}")
return accuracy
# --- CLI Entrypoint ---
def main():
parser = argparse.ArgumentParser()
parser.add_argument("model_id", help="HuggingFace model ID")
parser.add_argument(
"--stdout",
action="store_true",
help="Print accuracy to stdout instead of saving to a file"
)
args = parser.parse_args()
# For reproducibility
random.seed(42)
# Run evaluation
evaluate_multiplication(args.model_id, to_stdout=args.stdout)
if __name__ == "__main__":
main()