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)