frankenmoe-python-typescript / moe_orchestrator.py
punsaisuwan's picture
v2: Add streaming support + reasoning expert routing
13c1959 verified
Raw
History Blame Contribute Delete
26.2 kB
import re
import time
from datetime import datetime
from typing import Optional
from mlx_lm import load, generate, stream_generate
from mlx_lm.sample_utils import make_sampler, make_logits_processors
from moe_fixes import (
route_query, get_config, clean_or_flag_response, inject_persona,
limit_hedge_phrases, is_refusal, is_code_refusal, strip_translation_preamble,
detect_language_leakage_words, enforce_persona,
is_degenerate_query, DEGENERATE_FALLBACK_RESPONSE, get_greeting_quick_reply,
is_repetitive_response
)
# หมายเหตุสำคัญ: ไม่เปิด fix_mistral_regex เพราะพิสูจน์แล้วในโปรเจกต์นี้ว่า
# ทำให้เกิด Train-Inference Tokenizer Mismatch — LoRA Adapter จะตอบผิดเพี้ยน
# Repetition Penalty (Round 9): ป้องกัน Generation Collapse ที่พบใน
# Multi-turn Test ทั้ง Python Expert และ Reasoning Expert — ค่า 1.3 เป็น
# ค่ากลางที่ลด Loop ได้โดยไม่กระทบความลื่นไหลของภาษาไทยมากเกินไป
REPETITION_PENALTY = 1.15
REPETITION_CONTEXT_SIZE = 20
class MoEOrchestrator:
MAX_HISTORY_MESSAGES = 6
def __init__(self, base_model_path: str, adapter_paths: dict, log_file: str = "route_log.txt"):
self.log_file = log_file
print("[INFO] Loading base model (no adapter, ใช้สำหรับ Translation + Classification)...")
self.base_model, self.tokenizer = load(base_model_path)
self.experts = {}
for name, path in adapter_paths.items():
print(f"[INFO] Loading Expert: {name} ...")
model, _ = load(base_model_path, adapter_path=path)
self.experts[name] = model
print("[INFO] All experts loaded and cached in memory. Ready.\n")
def format_prompt(self, user_content: str, use_persona: bool = True, history: Optional[list] = None) -> str:
messages = []
if use_persona:
messages.append({"role": "system", "content": inject_persona("")})
if history:
trimmed_history = history[-self.MAX_HISTORY_MESSAGES:]
messages.extend(trimmed_history)
messages.append({"role": "user", "content": user_content})
return self.tokenizer.apply_chat_template(
messages,
tokenize=False,
add_generation_prompt=True
)
def check_truncation(self, text: str, max_tokens: int) -> bool:
token_count = len(self.tokenizer.encode(text))
return token_count >= max_tokens - 5
def llm_classify_route(self, query: str, fallback_expert: str = "reasoning") -> str:
classify_content = (
"Classify the following user query into exactly ONE category. "
"Reply with ONLY one word: python, typescript, or reasoning.\n\n"
f'Query: "{query}"\n\n'
"Category:"
)
prompt = self.format_prompt(classify_content, use_persona=False)
classify_sampler = make_sampler(temp=0.1, top_p=0.9)
raw_result = generate(
self.base_model, self.tokenizer, prompt=prompt,
max_tokens=10, sampler=classify_sampler
)
result = raw_result.strip().lower()
print(f"[CLASSIFIER-DEBUG] raw output: {raw_result!r}")
if "python" in result:
return "python"
elif "typescript" in result:
return "typescript"
elif "reasoning" in result:
return "reasoning"
else:
print(f"[WARNING] ⚠️ Classifier ตอบไม่ตรงหมวดใดเลย ('{result}') → ใช้ Weighted Score Fallback: {fallback_expert}")
return fallback_expert
def route(self, query: str, previous_expert: Optional[str] = None) -> tuple:
result = route_query(query, previous_expert=previous_expert)
expert = result["expert"]
method = result["method"]
scores = result["scores"]
print(f"[ROUTER-DEBUG] scores={scores} ambiguous={result['is_ambiguous']} method={method}")
if method in ("history_continuation", "history_continuation_debug"):
print(f"[INFO] Query ไม่มี Keyword ชัดเจน แต่มี previous_expert='{previous_expert}' → ใช้ต่อทันที (ข้าม Classifier)")
return expert, method
if expert is None or result["is_ambiguous"]:
code_scores = {"python": scores.get("python", 0.0), "typescript": scores.get("typescript", 0.0)}
has_code_signal = any(v > 0 for v in code_scores.values())
reasoning_score = scores.get("reasoning", 0.0)
if has_code_signal and reasoning_score == 0.0:
expert = max(code_scores, key=code_scores.get)
method = "weighted_tiebreak_code_only"
print(f"[INFO] ความกำกวมเป็นแค่ระหว่าง python/typescript (reasoning=0) → ใช้ Weighted Score ตัดสินตรงๆ: {expert}")
return expert, method
if any(v > 0 for v in scores.values()):
weighted_fallback = max(scores, key=scores.get)
else:
weighted_fallback = previous_expert if previous_expert else "reasoning"
print("[INFO] Weighted router ตัดสินใจไม่ได้จริง → ใช้ LLM Classifier ช่วย...")
expert = self.llm_classify_route(query, fallback_expert=weighted_fallback)
method = "llm_classifier"
return expert, method
def extract_code_blocks(self, text: str):
text = re.sub(r'\be\.g\.\b', 'for example', text, flags=re.IGNORECASE)
text = re.sub(r'\bi\.e\.\b', 'that is', text, flags=re.IGNORECASE)
code_blocks = re.findall(r"```[\s\S]*?```", text)
placeholder_text = text
for i, block in enumerate(code_blocks):
placeholder_text = placeholder_text.replace(block, f"[[CODE_BLOCK_{i}]]", 1)
return placeholder_text, code_blocks
def restore_code_blocks(self, translated_text: str, code_blocks: list):
for i, block in enumerate(code_blocks):
translated_text = translated_text.replace(f"[[CODE_BLOCK_{i}]]", block)
return translated_text
def _run_translation(self, text_only: str, max_tokens: int, minimal: bool = False) -> str:
if minimal:
translation_content = f"Translate this text into Thai:\n\n{text_only}"
else:
translation_content = (
"Translate the following text to Thai naturally.\n"
"Use 'ผม' as the first-person pronoun (never 'ฉัน' or 'ดิฉัน'). "
"End sentences with 'ครับ' where appropriate.\n"
"Keep [[CODE_BLOCK_N]] placeholders and all numbers unchanged.\n"
"Output ONLY the translated text, no extra commentary:\n\n"
f"{text_only}"
)
prompt = self.format_prompt(translation_content, use_persona=False)
translate_sampler = make_sampler(temp=0.3, top_p=0.9)
return generate(
self.base_model, self.tokenizer, prompt=prompt,
max_tokens=max_tokens, sampler=translate_sampler
)
def _has_translation_problem(self, text: str) -> bool:
if is_refusal(text):
return True
if detect_language_leakage_words(text):
return True
return False
def translate_to_thai(self, text: str, max_tokens: int = 800) -> tuple:
text_only, code_blocks = self.extract_code_blocks(text)
thai_text = self._run_translation(text_only, max_tokens, minimal=False)
if self._has_translation_problem(thai_text):
problem_reason = "Refusal" if is_refusal(thai_text) else "Language Leakage"
print(f"[WARNING] ⚠️ พบปัญหาการแปล ({problem_reason}) — Retry ครั้งเดียว (จำกัดไม่ให้ Retry ซ้ำ)...")
thai_text = self._run_translation(text_only, max_tokens, minimal=True)
thai_text = strip_translation_preamble(thai_text)
if self._has_translation_problem(thai_text):
print("[WARNING] ⚠️ Retry แล้วยังมีปัญหาอยู่ — ยอมรับผลลัพธ์นี้ไปก่อน (ไม่ Retry ซ้ำเพื่อคุม Latency)")
truncated = self.check_truncation(thai_text, max_tokens)
thai_text = limit_hedge_phrases(thai_text, max_allowed=1)
restored_text = self.restore_code_blocks(thai_text, code_blocks)
restored_text = enforce_persona(restored_text)
purity = clean_or_flag_response(restored_text)
if not purity["is_clean"]:
print(f"[WARNING] ⚠️ พบคำที่ไม่ใช่ภาษาไทยหลุดมาในคำแปล (ยอมรับผลลัพธ์นี้ไปพร้อม flag): {purity['flagged_words']}")
return purity["text"], truncated
def _generate_code_response(self, expert_model, prompt: str, config, effective_max_tokens: int) -> tuple:
logits_processors = make_logits_processors(
repetition_penalty=REPETITION_PENALTY,
repetition_context_size=REPETITION_CONTEXT_SIZE,
)
sampler = make_sampler(temp=config.temp, top_p=config.top_p)
response = generate(
expert_model, self.tokenizer, prompt=prompt,
max_tokens=effective_max_tokens,
sampler=sampler,
logits_processors=logits_processors,
)
has_refusal = is_code_refusal(response)
has_loop = is_repetitive_response(response)
if has_refusal or has_loop:
reason = "ปฏิเสธงานโดยไม่มีเหตุผลสมควร" if has_refusal else "ติด Repetition Loop (Generate ซ้ำวนไม่รู้จบ)"
print(f"[WARNING] ⚠️ Code Expert {reason} — Retry ครั้งเดียว (ป้องกัน Compounding Bias ลามไป History)...")
retry_sampler = make_sampler(temp=min(config.temp + 0.2, 1.0), top_p=config.top_p)
retry_logits_processors = make_logits_processors(
repetition_penalty=REPETITION_PENALTY,
repetition_context_size=REPETITION_CONTEXT_SIZE,
)
retry_response = generate(
expert_model, self.tokenizer, prompt=prompt,
max_tokens=effective_max_tokens,
sampler=retry_sampler,
logits_processors=retry_logits_processors,
)
if not (is_code_refusal(retry_response) or is_repetitive_response(retry_response)):
return retry_response, True
print("[WARNING] ⚠️ Retry แล้วยังมีปัญหาอยู่ — ยอมรับผลลัพธ์เดิมไปก่อน (ไม่ Retry ซ้ำเพื่อคุม Latency)")
return response, True
return response, False
# ---------- Streaming Support (Round 12: Stream Output) ----------
# หมายเหตุ: ใช้เฉพาะกับ run_interactive() ผ่าน query_stream() เท่านั้น
# query() เดิมด้านล่างไม่ถูกแก้ไข เพื่อไม่กระทบ test_multiturn.py
# ที่ต้องใช้ Return Value เป็น String เต็มก้อนสำหรับ Automated Check
def _generate_code_response_stream(self, expert_model, prompt: str, config, effective_max_tokens: int):
"""
Streaming version ของ _generate_code_response()
ใช้ Sliding Window Check ทุก ~50 Token แทนการตรวจทั้งก้อนหลัง Generate ครบ
เพราะ Streaming ทำให้ "ย้อนกลับไปแก้สิ่งที่ User เห็นไปแล้ว" ไม่ได้
ยอมรับ Trade-off ว่า User อาจเห็นการซ้ำ 1-2 รอบก่อนถูกตัด
(ดีกว่าเห็นซ้ำจนครบ max_tokens แบบเดิม)
"""
logits_processors = make_logits_processors(
repetition_penalty=REPETITION_PENALTY,
repetition_context_size=REPETITION_CONTEXT_SIZE,
)
sampler = make_sampler(temp=config.temp, top_p=config.top_p)
pending = "" # Buffer กัน enforce_persona() ตัดคำครึ่งที่ Boundary
check_window = ""
tokens_since_check = 0
CHECK_EVERY_N_TOKENS = 50
aborted = False
for gen_response in stream_generate(
expert_model, self.tokenizer, prompt=prompt,
max_tokens=effective_max_tokens,
sampler=sampler,
logits_processors=logits_processors,
):
delta = gen_response.text
pending += delta
check_window += delta
tokens_since_check += 1
# กันคำถูกตัดครึ่ง: เก็บ 10 ตัวท้ายสุดไว้ ค่อย Flush ส่วนที่ปลอดภัย
if len(pending) > 10:
safe_part = pending[:-10]
pending = pending[-10:]
yield enforce_persona(safe_part)
# Sliding Window Repetition Check ทุก ~50 Token
if tokens_since_check >= CHECK_EVERY_N_TOKENS:
if is_repetitive_response(check_window):
print("\n[WARNING] ⚠️ ตรวจพบ Repetition Loop กลาง Stream — ตัดจบทันที")
aborted = True
break
check_window = ""
tokens_since_check = 0
if pending:
yield enforce_persona(pending)
if aborted:
yield "\n\n[ตัดคำตอบเพื่อป้องกันการวนซ้ำ — ขออภัยในความไม่สะดวกครับ]"
self._last_stream_had_loop = aborted
def query_stream(self, user_query: str, max_tokens: Optional[int] = None,
history: Optional[list] = None, previous_expert: Optional[str] = None):
"""
เวอร์ชัน Streaming ของ query() — ใช้เฉพาะ run_interactive()
ไม่แทนที่ query() เดิม เพื่อไม่กระทบ test_multiturn.py ที่ต้องใช้ Return
เป็น String เต็มก้อนสำหรับ Automated Check
Yields: str (Chunk ของข้อความ) — หลัง Stream จบ self._last_expert_used
จะถูกตั้งค่าเหมือน query() ทุกประการ
"""
start_time = time.time()
if is_degenerate_query(user_query):
latency = time.time() - start_time
self.log_interaction(user_query, "none", "degenerate_fast_path", latency, False)
print(f"[INFO] ⚡ Fast-Path: Query ไม่มีเนื้อหาที่มีความหมาย → ตอบทันทีโดยไม่ผ่าน Model (Latency: {latency:.4f}s)")
self._last_expert_used = previous_expert
yield DEGENERATE_FALLBACK_RESPONSE
return
greeting_reply = get_greeting_quick_reply(user_query)
if greeting_reply is not None:
latency = time.time() - start_time
self.log_interaction(user_query, "none", "greeting_quick_reply", latency, False)
print(f"[INFO] ⚡ Quick-Reply: Greeting/Closing ล้วนๆ → ตอบทันทีโดยไม่ผ่าน Model (Latency: {latency:.4f}s)")
self._last_expert_used = previous_expert
yield greeting_reply
return
expert_name, method = self.route(user_query, previous_expert=previous_expert)
config = get_config(expert_name)
effective_max_tokens = max_tokens if max_tokens is not None else config.max_tokens
print(f"[ROUTER] เลือก Expert: {expert_name} (วิธี: {method}) | max_tokens={effective_max_tokens}")
expert_model = self.experts[expert_name]
prompt = self.format_prompt(user_query, history=history)
if expert_name == "reasoning":
# Reasoning ต้อง Translate หลัง Generate เสมอ — Stream ตรงๆจะโชว์
# ข้อความที่ยังไม่ใช่คำตอบสุดท้าย จึง Fallback ไปใช้ query() แบบ
# Blocking แล้ว Return ทั้งก้อนเป็น Chunk เดียว
final_response = self.query(
user_query, max_tokens=max_tokens,
history=history, previous_expert=previous_expert
)
yield final_response
return
try:
for chunk in self._generate_code_response_stream(
expert_model, prompt, config, effective_max_tokens
):
yield chunk
except Exception as e:
yield f"\n[ERROR] เกิดข้อผิดพลาดระหว่าง Generate: {str(e)}"
self._last_expert_used = expert_name
return
latency = time.time() - start_time
self.log_interaction(user_query, expert_name, method, latency, False)
print(f"\n[INFO] ใช้เวลา: {latency:.2f} วินาที")
self._last_expert_used = expert_name
def log_interaction(self, query: str, expert: str, method: str, latency: float, truncated: bool):
timestamp = datetime.now().strftime("%Y-%m-%d %H:%M:%S")
flag = " | ⚠️TRUNCATED" if truncated else ""
log_line = f"[{timestamp}] Expert={expert} | Method={method} | Latency={latency:.2f}s{flag} | Query={query}\n"
with open(self.log_file, "a", encoding="utf-8") as f:
f.write(log_line)
# ---------- Main Entry Point (Round 7: เพิ่ม Fast-Path Guard) ----------
def query(self, user_query: str, max_tokens: Optional[int] = None,
history: Optional[list] = None, previous_expert: Optional[str] = None) -> str:
"""
history: List ของ {"role": "user"/"assistant", "content": str} จาก Turn ก่อนหน้า
previous_expert: ชื่อ Expert ที่ใช้ใน Turn ก่อนหน้า
Fast-Path Guard (Round 7): ถ้า user_query ไม่มีตัวอักษรที่มีความหมาย
เพียงพอ (ว่าง/Emoji-only/Symbol-only) จะตอบด้วย Canned Response ทันที
โดยไม่แตะ Model เลย — ป้องกัน Hallucination และ Latency ที่ไม่จำเป็น
(พบจาก Adversarial Test ว่า Query ประเภทนี้กิน Latency 4-32 วินาที
และเสี่ยงหลุดคำตอบเพี้ยน เช่น พูดถึง "crab" ที่ไม่เกี่ยวข้อง)
"""
start_time = time.time()
if is_degenerate_query(user_query):
latency = time.time() - start_time
self.log_interaction(user_query, "none", "degenerate_fast_path", latency, False)
print(f"[INFO] ⚡ Fast-Path: Query ไม่มีเนื้อหาที่มีความหมาย → ตอบทันทีโดยไม่ผ่าน Model (Latency: {latency:.4f}s)")
self._last_expert_used = previous_expert
return DEGENERATE_FALLBACK_RESPONSE
greeting_reply = get_greeting_quick_reply(user_query)
if greeting_reply is not None:
latency = time.time() - start_time
self.log_interaction(user_query, "none", "greeting_quick_reply", latency, False)
print(f"[INFO] ⚡ Quick-Reply: Greeting/Closing ล้วนๆ → ตอบทันทีโดยไม่ผ่าน Model (Latency: {latency:.4f}s)")
self._last_expert_used = previous_expert
return greeting_reply
expert_name, method = self.route(user_query, previous_expert=previous_expert)
config = get_config(expert_name)
effective_max_tokens = max_tokens if max_tokens is not None else config.max_tokens
print(f"[ROUTER] เลือก Expert: {expert_name} (วิธี: {method}) | max_tokens={effective_max_tokens}")
expert_model = self.experts[expert_name]
prompt = self.format_prompt(user_query, history=history)
truncated = False
try:
if expert_name in ("python", "typescript"):
raw_response, _was_retried = self._generate_code_response(
expert_model, prompt, config, effective_max_tokens
)
else:
logits_processors = make_logits_processors(
repetition_penalty=REPETITION_PENALTY,
repetition_context_size=REPETITION_CONTEXT_SIZE,
)
sampler = make_sampler(temp=config.temp, top_p=config.top_p)
raw_response = generate(
expert_model, self.tokenizer, prompt=prompt,
max_tokens=effective_max_tokens,
sampler=sampler,
logits_processors=logits_processors,
)
if is_repetitive_response(raw_response):
print("[WARNING] ⚠️ Reasoning Expert ติด Repetition Loop — Retry ครั้งเดียว...")
retry_sampler = make_sampler(temp=min(config.temp + 0.2, 1.0), top_p=config.top_p)
retry_response = generate(
expert_model, self.tokenizer, prompt=prompt,
max_tokens=effective_max_tokens,
sampler=retry_sampler,
logits_processors=logits_processors,
)
if not is_repetitive_response(retry_response):
raw_response = retry_response
else:
print("[WARNING] ⚠️ Retry แล้วยังติด Loop อยู่ — ยอมรับผลลัพธ์เดิมไปก่อน (ไม่ Retry ซ้ำเพื่อคุม Latency)")
raw_truncated = self.check_truncation(raw_response, effective_max_tokens)
if expert_name == "reasoning":
print("[INFO] แปลผลลัพธ์เป็นไทย (Reasoning Expert)...")
final_response, trans_truncated = self.translate_to_thai(raw_response, max_tokens=effective_max_tokens)
truncated = raw_truncated or trans_truncated
else:
final_response = raw_response
truncated = raw_truncated
final_response = enforce_persona(final_response)
if truncated:
print("[WARNING] ⚠️ คำตอบอาจถูกตัดกลางคันเพราะใกล้ max_tokens — ควรเพิ่มค่านี้ถ้าเกิดบ่อย")
except Exception as e:
final_response = f"[ERROR] เกิดข้อผิดพลาดระหว่าง Generate: {str(e)}"
latency = time.time() - start_time
self.log_interaction(user_query, expert_name, method, latency, truncated)
print(f"[INFO] ใช้เวลา: {latency:.2f} วินาที")
self._last_expert_used = expert_name
return final_response
# ---------- Interactive Chat Loop ----------
def run_interactive(self):
print("=" * 60)
print("🤖 MoE Multi-Expert System — พิมพ์คำถามได้เลย")
print(" พิมพ์ 'exit' หรือ 'quit' เพื่อออกจากโปรแกรม")
print("=" * 60 + "\n")
history = []
last_expert = None
while True:
try:
user_input = input("คุณ: ").strip()
except (KeyboardInterrupt, EOFError):
print("\n[INFO] ออกจากโปรแกรม...")
break
if not user_input:
continue
if user_input.lower() in ["exit", "quit", "ออก"]:
print("[INFO] ออกจากโปรแกรม... แล้วพบกันใหม่ครับ 👋")
break
print()
print("🤖 ระบบ: ", end="", flush=True)
full_response = ""
for chunk in self.query_stream(user_input, history=history, previous_expert=last_expert):
print(chunk, end="", flush=True)
full_response += chunk
print("\n")
print("-" * 60 + "\n")
history.append({"role": "user", "content": user_input})
history.append({"role": "assistant", "content": full_response})
last_expert = self._last_expert_used
if __name__ == "__main__":
orchestrator = MoEOrchestrator(
base_model_path="./output-moe-mlx-4bit",
adapter_paths={
"python": "./adapters/expert-1-python",
"typescript": "./adapters/expert-2-typescript",
"reasoning": "./adapters/expert-3-reasoning",
}
)
orchestrator.run_interactive()