1import requests
2
3# Get a random Wikipedia article summary using their API
4def random_extract():
5 URL = "https://en.wikipedia.org/api/rest_v1/page/random/summary"
6 PARAMS = {}
7 r = requests.get(url = URL, params = PARAMS)
8 data = r.json()
9 return data['extract']
10
11# Format this as a prompt that would hopefully result in the model completing with a question
12def random_prompt():
13 e = random_extract()
14 return f"""### CONTEXT: {e} ### QUESTION:"""
15
16import torch
17from peft import AutoPeftModelForCausalLM
18from transformers import AutoTokenizer
19
20output_dir = "mcqgen_test"
21
22# load base LLM model and tokenizer
23model = AutoPeftModelForCausalLM.from_pretrained(
24 output_dir,
25 low_cpu_mem_usage=True,
26 torch_dtype=torch.float16,
27 load_in_4bit=True,
28)
29tokenizer = AutoTokenizer.from_pretrained(output_dir)
30
31# We can feed in a random context prompt and see what question the model comes up with:
32prompt = random_prompt()
33
34input_ids = tokenizer(prompt, return_tensors="pt", truncation=True).input_ids.cuda()
35# with torch.inference_mode():
36outputs = model.generate(input_ids=input_ids, max_new_tokens=100, do_sample=True, top_p=0.9,temperature=0.9)
37
38print(f"Prompt:\n{prompt}\n")
39print(f"Generated MCQ:\n### QUESTION:{tokenizer.batch_decode(outputs.detach().cpu().numpy(), skip_special_tokens=True)[0][len(prompt):]}")
40
41def process_outputs(outputs):
42 s = tokenizer.batch_decode(outputs.detach().cpu().numpy(), skip_special_tokens=True)[0]
43 split = s.split("### ")[1:][:7]
44 if len(split) != 7:
45 return None
46 # Check the starts
47 expected_starts = ['CONTEXT', 'QUESTION', 'A' , 'B', 'C', 'D', 'CORRECT']
48 for i, s in enumerate(split):
49 if not split[i].startswith(expected_starts[i]):
50 return None
51 return {
52 "context": split[0].replace("CONTEXT: ", ""),
53 "question": split[1].replace("QUESTION: ", ""),
54 "a": split[2].replace("A: ", ""),
55 "b": split[3].replace("B: ", ""),
56 "c": split[4].replace("C: ", ""),
57 "d": split[5].replace("D: ", ""),
58 "correct": split[6].replace("CORRECT: ", "")
59 }
60
61
62process_outputs(outputs) # A nice dictionary hopefully
63