test / multiq /cli.py
yousseftallal's picture
Feature: QuestionBank + game endpoints (/next, /prefill) + Groq provider
818b3aa
Raw History Blame Contribute Delete
9.96 kB
from __future__ import annotations
import argparse
import json
import logging
import sys
import time
from .config import load_settings
from .generator import QuestionGenerator
from .retriever import WikipediaRetriever
logging.basicConfig(level=logging.INFO, format="%(asctime)s [%(levelname)s] %(message)s")
def index_cmd(args):
"""Seed Wikipedia articles + build FAISS index."""
s = load_settings()
r = WikipediaRetriever(s)
titles = getattr(args, "titles", None) or (CORE_TOPICS if getattr(args, "auto_seed", False) else [])
if titles:
added = r.preload_pages(titles, lang="en")
logging.info("Seeded %d new articles", added)
n = r.build_index()
print(f"Indexed {n} documents into {s.index_path}")
CORE_TOPICS = [
"Photosynthesis", "Cell_(biology)", "DNA", "Gravity", "Earth_(planet)",
"Water_cycle", "Atom", "Chemical_reaction", "Evolution", "Solar_System",
"Electricity", "Light", "Human_anatomy", "Plant", "Ecosystem",
]
def seed_wikipedia_cmd(args):
"""Standalone: fetch & cache Wikipedia summaries via `wikipedia` lib."""
import wikipedia
wikipedia.wikipedia.USER_AGENT = "MultiQGenBot/0.1 (contact@multiqgen.github)"
lang = getattr(args, "lang", "en")
titles = args.titles or CORE_TOPICS
out_path = getattr(args, "out", None)
wikipedia.set_lang(lang)
s = load_settings()
out_path = out_path or s.sources_path
out_path.parent.mkdir(parents=True, exist_ok=True)
existing_urls = set()
if out_path.exists():
with open(out_path, encoding="utf-8") as fh:
for line in fh:
obj = json.loads(line)
existing_urls.add(obj.get("url"))
added = 0
with open(out_path, "a", encoding="utf-8") as fh:
for title in titles:
try:
page = wikipedia.page(title)
content_url = page.url
if content_url in existing_urls:
logging.info("Skip cached: %s", title)
continue
summary = wikipedia.summary(title, sentences=8)
if not summary:
continue
obj = {
"title": page.title,
"lang": lang,
"text": summary,
"url": content_url,
}
fh.write(json.dumps(obj, ensure_ascii=False) + "\n")
added += 1
logging.info("Cached: %s (%d chars)", title, len(summary))
time.sleep(1) # be polite
except Exception as e:
logging.warning("Failed '%s': %s", title, type(e).__name__)
print(f"Seeded {added} new articles into {out_path}")
def serve_cmd(args):
from .api import create_app
app = create_app()
app.run(host=getattr(args, "host", "0.0.0.0"), port=args.port, debug=False)
def finetune_cmd(args):
"""
Minimal LoRA finetuning script:
- uses datasets.load_dataset('cais/mmlu') + ('openai/gsm8k') to
produce a merged Q/A training set in Arabic & English
- writes LoRA weights to models/
"""
from datasets import concatenate_datasets, load_dataset
from peft import LoraConfig, get_peft_model
from transformers import (
AutoModelForCausalLM,
AutoTokenizer,
Trainer,
TrainingArguments,
)
out_dir = getattr(args, "out", "models/qwen-lora")
epochs = getattr(args, "epochs", 2)
s = load_settings()
en_mmlu = load_dataset("cais/mmlu", "all", split="train").select(range(5000))
en_gsm = load_dataset("openai/gsm8k", "main", split="train").select(range(3000))
def to_qa(ex):
return {"text": f"Q: {ex['question']}\nA: {ex['answer'].split('####')[-1].strip()}"}
mmlu_qa = en_mmlu.map(lambda x: {"question": x["question"], "answer": f"{x['choices'][x['answer']]}"})
gsm_qa = en_gsm.map(to_qa)
# NOTE: Arabic translation via NLLB not implemented in PoC; using English subset
combined = concatenate_datasets([mmlu_qa, gsm_qa])
tokenizer = AutoTokenizer.from_pretrained(s.llm_name)
tokenizer.pad_token = tokenizer.eos_token
model = AutoModelForCausalLM.from_pretrained(s.llm_name, torch_dtype="auto", device_map="auto")
lora = LoraConfig(
r=16,
lora_alpha=32,
target_modules=["q_proj", "k_proj", "v_proj", "o_proj"],
lora_dropout=0.05,
bias="none",
task_type="CAUSAL_LM",
)
model = get_peft_model(model, lora)
def tokenize_fn(x):
return tokenizer(x["text"], truncation=True, padding="max_length", max_length=512)
tokenized = combined.map(tokenize_fn, remove_columns=["text", "question", "answer", "choices", "subject", "extract"])
args_ = TrainingArguments(
output_dir=out_dir,
per_device_train_batch_size=4,
gradient_accumulation_steps=4,
num_train_epochs=epochs,
fp16=True,
logging_steps=100,
save_steps=500,
save_total_limit=2,
report_to="none",
)
trainer = Trainer(
model=model,
args=args_,
train_dataset=tokenized,
)
trainer.train()
model.save_pretrained(out_dir)
print("Fine-tuning complete")
def eval_cmd(args):
"""Evaluate accuracy + difficulty calibration."""
import os
s = load_settings()
# use smoke model for quick eval unless overridden
smoke = os.getenv("MQG_SMOKE", "1") in ("1", "true", "True")
gen = QuestionGenerator(s, smoke=smoke)
sample_topics = {
"easy": ["Photosynthesis", "Water cycle", "Gravity basics"],
"medium": ["Newton's laws", "Chemical bonding", "Egyptian civilization"],
"hard": ["Quantum tunneling", "Gödel's theorems", "Mitochondrial DNA"],
}
stats = {"easy": [], "medium": [], "hard": []}
for diff, topics in sample_topics.items():
for t in topics:
try:
q = gen.generate(topic=t, lang="en", difficulty=diff)
stats[diff].append(q.to_dict())
print(json.dumps(q.to_dict(), indent=2, ensure_ascii=False))
except Exception as e:
print({"error": str(e)}, file=sys.stderr)
print("Evaluation samples:", len(stats["hard"]), "hard", len(stats["medium"]), "medium", len(stats["easy"]), "easy")
def prefill_cmd(args):
"""Pre-generate a question bank (offline): topics x langs x difficulties x types.
Generates questions ahead of time and stores them in data/questions.jsonl so
the game can draw instantly with zero LLM latency / rate limits at runtime.
"""
from .bank import QuestionBank
s = load_settings()
bank = QuestionBank(s.bank_path)
gen = QuestionGenerator(s, smoke=args.smoke)
topics = args.topics or CORE_TOPICS
langs = args.langs
difficulties = args.difficulties
types = args.types
print(f"Bank file: {s.bank_path} (existing: {bank.stats()['total']})")
generated = 0
failed = 0
for topic in topics:
for lang in langs:
for diff in difficulties:
for qtype in types:
made = 0
attempts = 0
while made < args.per_cell and attempts < args.per_cell * 4:
attempts += 1
try:
q = gen.generate(topic=topic, lang=lang, difficulty=diff, qtype=qtype)
except Exception as e:
logging.error("Failed %s/%s/%s: %s", topic, lang, diff, str(e)[:200])
failed += 1
break
if bank.add(q.to_dict()):
made += 1
generated += 1
else:
attempts -= 1 # duplicate id → retry without burning an attempt
if made < args.per_cell:
logging.warning("Short on %s/%s/%s: %d/%d", topic, lang, diff, made, args.per_cell)
print(f"Done. Generated {generated} new questions. Bank total: {bank.stats()['total']}. Failures: {failed}")
if __name__ == "__main__":
parser = argparse.ArgumentParser(prog="mqg")
subparser = parser.add_subparsers(dest="command")
p_index = subparser.add_parser("index", help="Build FAISS index from cached summaries")
p_index.add_argument("--langs", nargs="*")
p_index.set_defaults(func=index_cmd)
p_seed = subparser.add_parser("seed-wikipedia", help="Fetch Wikipedia summaries")
p_seed.add_argument("--lang", default="en")
p_seed.add_argument("--titles", nargs="*", default=None)
p_seed.add_argument("--out", default=None)
p_seed.set_defaults(func=seed_wikipedia_cmd)
p_serve = subparser.add_parser("serve")
p_serve.add_argument("--host", default="0.0.0.0")
p_serve.add_argument("--port", type=int, default=5000)
p_serve.set_defaults(func=serve_cmd)
p_ft = subparser.add_parser("finetune")
p_ft.add_argument("--out", default="models/qwen-lora")
p_ft.add_argument("--epochs", type=int, default=2)
p_ft.set_defaults(func=finetune_cmd)
subparser.add_parser("eval").set_defaults(func=eval_cmd)
p_prefill = subparser.add_parser("prefill", help="Pre-generate question bank")
p_prefill.add_argument("--topics", nargs="*", default=None)
p_prefill.add_argument("--langs", nargs="*", default=["en", "ar"])
p_prefill.add_argument("--difficulties", nargs="*", default=["easy", "medium", "hard"])
p_prefill.add_argument("--types", nargs="*", default=["multiple_choice"])
p_prefill.add_argument("--per-cell", type=int, default=5)
p_prefill.add_argument("--smoke", action="store_true", help="use the small local model")
p_prefill.set_defaults(func=prefill_cmd)
args = parser.parse_args()
if not hasattr(args, "func"):
parser.print_help()
sys.exit(0)
args.func(args)