hotdogs commited on
Commit
df441e9
Β·
verified Β·
1 Parent(s): 1ca083b

Upload simple_router.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. simple_router.py +114 -0
simple_router.py ADDED
@@ -0,0 +1,114 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Simple Router β€” prompt classification + expert routing.
2
+ Classifies coding/math/chat then loads the right LoRA expert.
3
+ """
4
+ import torch
5
+ import re
6
+ from transformers import AutoModelForCausalLM, AutoTokenizer
7
+
8
+ # ─── Classification ───
9
+ def classify_prompt(text: str) -> str:
10
+ """Classify prompt into domain using keyword matching."""
11
+ text_lower = text.lower()
12
+
13
+ coding_kw = [
14
+ 'def ', 'function', 'python', 'code', 'bug', 'debug', 'api',
15
+ 'import', 'class ', 'algorithm', 'implement', 'compile', 'syntax',
16
+ 'javascript', 'html', 'css', 'sql', 'bash', 'git', 'docker',
17
+ 'write a', 'program', 'script', 'loop', 'array', 'list', 'dict',
18
+ ]
19
+ math_kw = [
20
+ 'solve', 'equation', 'derivative', 'integral', 'matrix', 'eigen',
21
+ 'theorem', 'proof', 'sqrt', 'log', 'sin', 'cos', 'tan', 'sum',
22
+ 'probability', 'statistic', 'graph', 'vector', 'polynomial',
23
+ 'x =', 'x=', 'y =', 'calculate', 'compute', 'find the',
24
+ ]
25
+
26
+ coding_score = sum(1 for kw in coding_kw if kw in text_lower)
27
+ math_score = sum(1 for kw in math_kw if kw in text_lower)
28
+
29
+ if coding_score > 0 and coding_score >= math_score:
30
+ return 'coding'
31
+ elif math_score > 0 and math_score > coding_score:
32
+ return 'math'
33
+ else:
34
+ return 'chat'
35
+
36
+ # ─── Router ───
37
+ class ExpertRouter:
38
+ def __init__(self, base_model_name: str = 'unsloth/Qwen2.5-1.5B-Instruct'):
39
+ self.base_name = base_model_name
40
+ self.device = 'cuda' if torch.cuda.is_available() else 'cpu'
41
+
42
+ # Load base model once
43
+ print(f'Loading base model: {base_model_name}...')
44
+ self.base = AutoModelForCausalLM.from_pretrained(
45
+ base_model_name,
46
+ torch_dtype=torch.bfloat16,
47
+ device_map='auto',
48
+ )
49
+ self.tokenizer = AutoTokenizer.from_pretrained(base_model_name)
50
+ if self.tokenizer.pad_token is None:
51
+ self.tokenizer.pad_token = self.tokenizer.eos_token
52
+
53
+ # Cache for loaded experts
54
+ self._experts = {}
55
+ self._current = None
56
+
57
+ def _load_expert(self, domain: str):
58
+ """Load LoRA adapter for a domain."""
59
+ from peft import PeftModel
60
+ import copy
61
+
62
+ if domain in self._experts:
63
+ return self._experts[domain]
64
+
65
+ print(f' Loading expert: {domain}...')
66
+ # Load fresh base + adapter each time (merge_and_unload corrupts base)
67
+ base_fresh = AutoModelForCausalLM.from_pretrained(
68
+ self.base_name,
69
+ torch_dtype=torch.bfloat16,
70
+ device_map='auto',
71
+ )
72
+ model = PeftModel.from_pretrained(
73
+ base_fresh,
74
+ 'hotdogs/frankenmoe',
75
+ subfolder=domain,
76
+ torch_dtype=torch.bfloat16,
77
+ )
78
+ model = model.merge_and_unload()
79
+ self._experts[domain] = model
80
+ return model
81
+
82
+ def generate(self, prompt: str, max_tokens: int = 128) -> tuple:
83
+ """Classify + generate response."""
84
+ domain = classify_prompt(prompt)
85
+ model = self._load_expert(domain)
86
+
87
+ inp = self.tokenizer(prompt, return_tensors='pt').to(self.device)
88
+ with torch.no_grad():
89
+ out = model.generate(
90
+ **inp,
91
+ max_new_tokens=max_tokens,
92
+ do_sample=True,
93
+ temperature=0.7,
94
+ top_p=0.9,
95
+ )
96
+ text = self.tokenizer.decode(out[0], skip_special_tokens=True)
97
+ return domain, text
98
+
99
+ # ─── Main ───
100
+ if __name__ == '__main__':
101
+ router = ExpertRouter()
102
+
103
+ tests = [
104
+ 'Write a Python function to reverse a linked list',
105
+ 'Solve the quadratic equation 2x^2 - 4x + 1 = 0',
106
+ 'What is the capital of Thailand?',
107
+ ]
108
+
109
+ for prompt in tests:
110
+ domain, response = router.generate(prompt)
111
+ print(f'\n{"="*50}')
112
+ print(f'[ROUTE: {domain}]')
113
+ print(f'PROMPT: {prompt}')
114
+ print(f'OUTPUT: {response[-300:]}')