frankenmoe / simple_router.py
hotdogs's picture
Upload simple_router.py with huggingface_hub
df441e9 verified
Raw
History Blame Contribute Delete
4.05 kB
"""Simple Router β€” prompt classification + expert routing.
Classifies coding/math/chat then loads the right LoRA expert.
"""
import torch
import re
from transformers import AutoModelForCausalLM, AutoTokenizer
# ─── Classification ───
def classify_prompt(text: str) -> str:
"""Classify prompt into domain using keyword matching."""
text_lower = text.lower()
coding_kw = [
'def ', 'function', 'python', 'code', 'bug', 'debug', 'api',
'import', 'class ', 'algorithm', 'implement', 'compile', 'syntax',
'javascript', 'html', 'css', 'sql', 'bash', 'git', 'docker',
'write a', 'program', 'script', 'loop', 'array', 'list', 'dict',
]
math_kw = [
'solve', 'equation', 'derivative', 'integral', 'matrix', 'eigen',
'theorem', 'proof', 'sqrt', 'log', 'sin', 'cos', 'tan', 'sum',
'probability', 'statistic', 'graph', 'vector', 'polynomial',
'x =', 'x=', 'y =', 'calculate', 'compute', 'find the',
]
coding_score = sum(1 for kw in coding_kw if kw in text_lower)
math_score = sum(1 for kw in math_kw if kw in text_lower)
if coding_score > 0 and coding_score >= math_score:
return 'coding'
elif math_score > 0 and math_score > coding_score:
return 'math'
else:
return 'chat'
# ─── Router ───
class ExpertRouter:
def __init__(self, base_model_name: str = 'unsloth/Qwen2.5-1.5B-Instruct'):
self.base_name = base_model_name
self.device = 'cuda' if torch.cuda.is_available() else 'cpu'
# Load base model once
print(f'Loading base model: {base_model_name}...')
self.base = AutoModelForCausalLM.from_pretrained(
base_model_name,
torch_dtype=torch.bfloat16,
device_map='auto',
)
self.tokenizer = AutoTokenizer.from_pretrained(base_model_name)
if self.tokenizer.pad_token is None:
self.tokenizer.pad_token = self.tokenizer.eos_token
# Cache for loaded experts
self._experts = {}
self._current = None
def _load_expert(self, domain: str):
"""Load LoRA adapter for a domain."""
from peft import PeftModel
import copy
if domain in self._experts:
return self._experts[domain]
print(f' Loading expert: {domain}...')
# Load fresh base + adapter each time (merge_and_unload corrupts base)
base_fresh = AutoModelForCausalLM.from_pretrained(
self.base_name,
torch_dtype=torch.bfloat16,
device_map='auto',
)
model = PeftModel.from_pretrained(
base_fresh,
'hotdogs/frankenmoe',
subfolder=domain,
torch_dtype=torch.bfloat16,
)
model = model.merge_and_unload()
self._experts[domain] = model
return model
def generate(self, prompt: str, max_tokens: int = 128) -> tuple:
"""Classify + generate response."""
domain = classify_prompt(prompt)
model = self._load_expert(domain)
inp = self.tokenizer(prompt, return_tensors='pt').to(self.device)
with torch.no_grad():
out = model.generate(
**inp,
max_new_tokens=max_tokens,
do_sample=True,
temperature=0.7,
top_p=0.9,
)
text = self.tokenizer.decode(out[0], skip_special_tokens=True)
return domain, text
# ─── Main ───
if __name__ == '__main__':
router = ExpertRouter()
tests = [
'Write a Python function to reverse a linked list',
'Solve the quadratic equation 2x^2 - 4x + 1 = 0',
'What is the capital of Thailand?',
]
for prompt in tests:
domain, response = router.generate(prompt)
print(f'\n{"="*50}')
print(f'[ROUTE: {domain}]')
print(f'PROMPT: {prompt}')
print(f'OUTPUT: {response[-300:]}')