Download chat.py from FannyFa/FannyFa-Model-V1: direct link, hf CLI and curl.
- Browser
- Download file 40.8 kB
-
https://huggingface.co/FannyFa/FannyFa-Model-V1/resolve/main/chat.py
- Command line
-
hf download hf://FannyFa/FannyFa-Model-V1/chat.py
-
curl -L -o chat.py https://huggingface.co/FannyFa/FannyFa-Model-V1/resolve/main/chat.py
40.8 kB
| import os | |
| import re | |
| import json | |
| import time | |
| import random | |
| from datetime import datetime | |
| from pathlib import Path | |
| from collections import deque | |
| from typing import Optional, Dict, Any, List, Tuple | |
| import torch | |
| from transformers import AutoTokenizer, AutoModelForCausalLM, TextIteratorStreamer | |
| from threading import Thread | |
| from rich.markup import escape as rich_escape | |
| from rich.panel import Panel | |
| from rich.prompt import Prompt, Confirm | |
| from rich.table import Table | |
| from rich import box | |
| from config import MODEL_DIR | |
| from models import ModelManager | |
| from utils import console, Theme, error_logger, debug_logger, get_device | |
| # --------------------------------------------------------------------------- | |
| # Command aliases (typo-tolerant) | |
| # --------------------------------------------------------------------------- | |
| _COMMAND_ALIASES = { | |
| "xlear": "clear", "claer": "clear", "clea": "clear", | |
| "exi": "exit", "ext": "exit", "quitt": "exit", | |
| "hlp": "help", "hep": "help", | |
| "hist": "history", "his": "history", | |
| "rg": "regen", "regen": "regen", | |
| "u": "undo", | |
| "s": "stats", "st": "stats", | |
| "e": "exit", | |
| } | |
| _DEFAULT_PERSONA = "Anda adalah asisten AI yang ramah dan informatif." | |
| # Boilerplate umum dari model yang wajib dipotong | |
| _BOILERPLATE_MARKERS = [ | |
| "semoga membantu", | |
| "semoga informasinya jelas", | |
| "saya siap membantu lebih lanjut", | |
| "ada lagi yang ingin ditanyakan", | |
| "ada lagi yang bisa dibantu", | |
| "kalau ada pertanyaan lain", | |
| "apakah ada yang ingin didiskusikan", | |
| "jangan ragu untuk bertanya", | |
| "terima kasih sudah bertanya", | |
| "semoga harimu menyenangkan", | |
| "bisa ditanyakan lebih detail", | |
| "saya di sini untuk membantu", | |
| "mari jaga percakapan", | |
| "semoga kita bisa bicara", | |
| "saya mohon maaf atas ketidaknyamanan", | |
| "mari fokus pada solusi", | |
| "silakan ajukan pertanyaan yang sopan", | |
| ] | |
| # Frasa sampah yang menandakan output ngelantur (typo atau halusinasi) | |
| _JUNK_PHRASES = [ | |
| "seminta", "semohon", "semiap", "sembuh", "semangat membantu", | |
| "saya di rumah", "hati-hati di jalan", "jangan lupa pemanasan", | |
| "jangan lupa pendinginan", "saya siap belajar", | |
| "saya catat", "hari baik", "di sana untuk membantu", | |
| "baik bicara lebih baik", "baik yang baik", | |
| "bicara lebih positif", "mari berdiskusi", | |
| "silakan bicara dengan tenang", "mari siap membantu", | |
| ] | |
| # Kata kasar user yang sering bikin model output aneh | |
| _RUDE_WORDS = { | |
| "anjing", "bangsat", "kontol", "memek", "goblok", "tolol", "idiot", | |
| "bego", "bodoh", "kampret", "sialan", "bajingan", "keparat", | |
| } | |
| # Fallback responses | |
| _FALLBACKS = { | |
| "greeting": "Halo! Ada yang bisa saya bantu?", | |
| "howareyou": "Saya baik, terima kasih. Ada yang bisa dibantu?", | |
| "thanks": "Sama-sama! Senang bisa membantu.", | |
| "bye": "Sampai jumpa! Senang bisa membantu.", | |
| "rude": "Yuk, kita jaga percakapan tetap sopan ya. Ada yang bisa saya bantu?", | |
| "unknown": "Maaf, saya belum bisa menjawab itu dengan baik. Bisa dicoba dengan pertanyaan lain?", | |
| "short": "Bisa dijelaskan sedikit lebih detail?", | |
| } | |
| # Keyword untuk deteksi intent sederhana | |
| _GREETINGS = {"halo", "hai", "hi", "hello", "woi", "hei", "assalamualaikum", "pagi", "siang", "sore", "malam"} | |
| _THANKS = {"makasih", "terima kasih", "thanks", "thank you", "thx", "tq"} | |
| _BYES = {"bye", "dadah", "sampai jumpa", "selamat tinggal", "goodbye"} | |
| _SESSION_DIR = Path("/content/data/chat_sessions") | |
| class EnhancedChatModule: | |
| def __init__(self): | |
| self.model_manager = ModelManager() | |
| self.device = get_device() | |
| self.tokenizer = None | |
| self.model = None | |
| self.conversation_history = deque(maxlen=50) | |
| self.max_context_tokens = 768 | |
| self.system_prompt = _DEFAULT_PERSONA | |
| self.is_peft_model = False | |
| # Konfigurasi generation yang lebih ketat untuk model kecil | |
| self.gen_config = { | |
| "temperature": 0.75, | |
| "top_p": 0.9, | |
| "top_k": 40, | |
| "max_new_tokens": 60, # pendek = lebih aman | |
| "repetition_penalty": 1.3, # lebih tinggi dari sebelumnya | |
| "no_repeat_ngram_size": 4, # lebih agresif | |
| "streaming": False, | |
| # Best-of-N sampling | |
| "num_attempts": 3, | |
| "min_acceptable_score": 0.45, | |
| } | |
| self.performance_stats = { | |
| "total_generations": 0, | |
| "avg_time_ms": 0.0, | |
| "total_tokens_generated": 0, | |
| "total_tokens_processed": 0, | |
| "max_response_time": 0.0, | |
| "min_response_time": float('inf'), | |
| "last_response": "", | |
| "last_user_input": "", | |
| "fallbacks_triggered": 0, | |
| "regenerations": 0, | |
| } | |
| _SESSION_DIR.mkdir(parents=True, exist_ok=True) | |
| # ================================================================== | |
| # Main loop | |
| # ================================================================== | |
| def run(self) -> None: | |
| current_model = self.model_manager.get_current_model() | |
| if not current_model: | |
| console.print(Theme.error(" Belum ada model. Training dulu!")) | |
| return | |
| model_path = os.path.join(MODEL_DIR, current_model) | |
| if not os.path.exists(model_path): | |
| console.print(Theme.error(" Model tidak ditemukan!")) | |
| return | |
| console.print(Theme.info(f"Memuat model: {current_model}...")) | |
| try: | |
| self._load_model(model_path) | |
| self.model.eval() | |
| console.print(Theme.success(" Model siap!")) | |
| except Exception as e: | |
| console.print(Theme.error(f"Gagal memuat model: {e}")) | |
| error_logger.error(f"Model loading error: {e}") | |
| return | |
| # Sync config dari ModelManager (temperature dll) | |
| self._sync_gen_config_from_manager() | |
| console.print(Panel( | |
| "[cyan] ENHANCED CHAT v2.2[/cyan]\n" | |
| "[white]Ketik 'help' untuk bantuan, 'exit' keluar[/white]\n" | |
| f"[dim]Model: {current_model} | Device: {self.device.upper()} | " | |
| f"Best-of-{self.gen_config['num_attempts']} | " | |
| f"Streaming: {'ON' if self.gen_config['streaming'] else 'OFF'}[/dim]", | |
| width=72, style="cyan" | |
| )) | |
| while True: | |
| try: | |
| user_input = Prompt.ask("[cyan]You").strip() | |
| if not user_input: | |
| continue | |
| lower = user_input.lower().strip() | |
| if lower in _COMMAND_ALIASES: | |
| lower = _COMMAND_ALIASES[lower] | |
| if self._handle_command(lower, user_input): | |
| continue | |
| self._generate_and_display(user_input) | |
| except KeyboardInterrupt: | |
| console.print("\n" + Theme.warning("Chat dihentikan")) | |
| break | |
| except Exception as e: | |
| console.print(Theme.error(f"Error: {rich_escape(str(e))}")) | |
| error_logger.error(f"Chat error: {e}") | |
| if "CUDA" in str(e) or "out of memory" in str(e).lower(): | |
| try: | |
| self.model = self.model.to('cpu') | |
| torch.cuda.empty_cache() | |
| self.device = 'cpu' | |
| console.print(Theme.success(" Moved to CPU")) | |
| except Exception: | |
| break | |
| # ================================================================== | |
| # Command handler | |
| # ================================================================== | |
| def _handle_command(self, lower: str, raw: str) -> bool: | |
| if lower in ('exit', 'q'): | |
| console.print(Theme.success("Terima kasih! ")) | |
| raise SystemExit | |
| if lower == 'clear': | |
| self.conversation_history.clear() | |
| console.print(Theme.success(" Riwayat dihapus")) | |
| return True | |
| if lower in ('history', 'hist'): | |
| self._show_history() | |
| return True | |
| if lower == 'stats': | |
| self._show_stats() | |
| return True | |
| if lower == 'help': | |
| self._show_help() | |
| return True | |
| if lower in ('undo', 'u'): | |
| if self.conversation_history: | |
| removed = self.conversation_history.pop() | |
| console.print(Theme.success(f"↩ Undo: hapus '{removed['user'][:40]}...'")) | |
| else: | |
| console.print(Theme.warning("Belum ada riwayat")) | |
| return True | |
| if lower in ('regen', 'rg'): | |
| last_user = self.performance_stats.get("last_user_input") | |
| if not last_user: | |
| console.print(Theme.warning("Belum ada percakapan")) | |
| return True | |
| if self.conversation_history: | |
| self.conversation_history.pop() | |
| self.performance_stats["regenerations"] += 1 | |
| console.print(Theme.info(" Regenerate...")) | |
| self._generate_and_display(last_user) | |
| return True | |
| if lower == 'save': | |
| self._save_session() | |
| return True | |
| if lower == 'sessions': | |
| self._list_sessions() | |
| return True | |
| if lower.startswith('load'): | |
| parts = raw.split(maxsplit=1) | |
| if len(parts) < 2: | |
| self._list_sessions() | |
| return True | |
| try: | |
| idx = int(parts[1]) | |
| self._load_session(idx) | |
| except ValueError: | |
| console.print(Theme.error("Format: load <nomor>")) | |
| return True | |
| if lower == 'export': | |
| self._export_chat() | |
| return True | |
| if lower.startswith('stream '): | |
| val = lower.split(maxsplit=1)[1] | |
| if val in ("on", "true", "1", "yes"): | |
| self.gen_config["streaming"] = True | |
| console.print(Theme.success(" Streaming: ON")) | |
| elif val in ("off", "false", "0", "no"): | |
| self.gen_config["streaming"] = False | |
| console.print(Theme.success(" Streaming: OFF")) | |
| else: | |
| console.print(Theme.error("Format: stream on|off")) | |
| return True | |
| if lower.startswith('temp '): | |
| try: | |
| val = float(lower.split(maxsplit=1)[1]) | |
| if 0.01 <= val <= 2.0: | |
| self.gen_config["temperature"] = val | |
| console.print(Theme.success(f" Temperature → {val}")) | |
| else: | |
| console.print(Theme.warning("Nilai harus 0.01-2.0")) | |
| except (ValueError, IndexError): | |
| console.print(Theme.error("Format: temp <0.01-2.0>")) | |
| return True | |
| if lower.startswith('topp '): | |
| try: | |
| val = float(lower.split(maxsplit=1)[1]) | |
| if 0.0 <= val <= 1.0: | |
| self.gen_config["top_p"] = val | |
| console.print(Theme.success(f" Top-p → {val}")) | |
| else: | |
| console.print(Theme.warning("Nilai harus 0.0-1.0")) | |
| except (ValueError, IndexError): | |
| console.print(Theme.error("Format: topp <0.0-1.0>")) | |
| return True | |
| if lower.startswith('maxtok '): | |
| try: | |
| val = int(lower.split(maxsplit=1)[1]) | |
| if 10 <= val <= 500: | |
| self.gen_config["max_new_tokens"] = val | |
| console.print(Theme.success(f" Max new tokens → {val}")) | |
| else: | |
| console.print(Theme.warning("Nilai harus 10-500")) | |
| except (ValueError, IndexError): | |
| console.print(Theme.error("Format: maxtok <10-500>")) | |
| return True | |
| if lower.startswith('attempts '): | |
| try: | |
| val = int(lower.split(maxsplit=1)[1]) | |
| if 1 <= val <= 10: | |
| self.gen_config["num_attempts"] = val | |
| console.print(Theme.success(f" Best-of-N → {val}")) | |
| else: | |
| console.print(Theme.warning("Nilai harus 1-10")) | |
| except (ValueError, IndexError): | |
| console.print(Theme.error("Format: attempts <1-10>")) | |
| return True | |
| if lower.startswith('persona '): | |
| new_persona = raw[8:].strip() | |
| if new_persona.lower() == "reset": | |
| self.system_prompt = _DEFAULT_PERSONA | |
| console.print(Theme.success(" Persona di-reset ke default")) | |
| elif new_persona: | |
| self.system_prompt = new_persona | |
| console.print(Theme.success(f" Persona di-set ({len(new_persona)} chars)")) | |
| return True | |
| if lower == 'persona': | |
| console.print(Panel( | |
| f"[cyan]Current persona:[/cyan]\n{rich_escape(self.system_prompt)}", | |
| title="PERSONA", style="cyan" | |
| )) | |
| return True | |
| return False | |
| # ================================================================== | |
| # Generate & display (BEST-OF-N) | |
| # ================================================================== | |
| def _generate_and_display(self, user_input: str) -> None: | |
| start = time.perf_counter() | |
| # ----- Fast path: fallback untuk greeting / thanks / bye / rude ----- | |
| fast = self._fast_intent_response(user_input) | |
| if fast is not None: | |
| response = fast | |
| elif self.gen_config.get("streaming", False): | |
| response = self._generate_streaming(user_input) | |
| else: | |
| # Best-of-N | |
| response = self._generate_best_of_n(user_input) | |
| elapsed_ms = (time.perf_counter() - start) * 1000 | |
| if not self.gen_config.get("streaming", False): | |
| console.print(f"[magenta]AI[/]: {rich_escape(response)}") | |
| console.print(Theme.dim(f" {elapsed_ms:.0f}ms | {len(response.split())} words")) | |
| console.print() | |
| self._update_stats(elapsed_ms, response) | |
| self.performance_stats["last_user_input"] = user_input | |
| self.performance_stats["last_response"] = response | |
| self.conversation_history.append({ | |
| 'user': user_input, | |
| 'ai': response, | |
| 'timestamp': time.time(), | |
| }) | |
| # ================================================================== | |
| # Fast intent — deteksi sederhana tanpa model | |
| # ================================================================== | |
| def _fast_intent_response(self, user_input: str) -> Optional[str]: | |
| """Kalau intent jelas, jawab langsung. Model tidak dipanggil.""" | |
| lower = user_input.lower().strip() | |
| words = set(re.findall(r'\w+', lower)) | |
| # Kasar → jangan kasih ke model | |
| if words & _RUDE_WORDS: | |
| return _FALLBACKS["rude"] | |
| # Very short input (< 3 chars) | |
| if len(lower) < 3: | |
| return _FALLBACKS["short"] | |
| # Greeting | |
| if len(words) <= 3 and (words & _GREETINGS): | |
| return _FALLBACKS["greeting"] | |
| # Thanks | |
| if words & _THANKS: | |
| return _FALLBACKS["thanks"] | |
| # Bye | |
| if any(b in lower for b in _BYES): | |
| return _FALLBACKS["bye"] | |
| # "Apa kabar" | |
| if "kabar" in lower and len(words) <= 4: | |
| return _FALLBACKS["howareyou"] | |
| return None | |
| # ================================================================== | |
| # Best-of-N generation | |
| # ================================================================== | |
| def _generate_best_of_n(self, user_input: str) -> str: | |
| """Generate N kandidat, pilih skor tertinggi.""" | |
| n = max(1, self.gen_config.get("num_attempts", 3)) | |
| candidates: List[Tuple[float, str]] = [] | |
| context = self._build_context(user_input) | |
| for i in range(n): | |
| try: | |
| raw = self._generate_raw(context, attempt_index=i) | |
| cleaned = self._clean_response(raw, user_input) | |
| score = self._score_response(cleaned, user_input) | |
| candidates.append((score, cleaned)) | |
| except Exception as e: | |
| debug_logger.debug(f"Attempt {i} failed: {e}") | |
| continue | |
| if not candidates: | |
| self.performance_stats["fallbacks_triggered"] += 1 | |
| return _FALLBACKS["unknown"] | |
| candidates.sort(key=lambda x: x[0], reverse=True) | |
| best_score, best_response = candidates[0] | |
| min_score = self.gen_config.get("min_acceptable_score", 0.4) | |
| if best_score < min_score: | |
| self.performance_stats["fallbacks_triggered"] += 1 | |
| return _FALLBACKS["unknown"] | |
| return best_response | |
| def _generate_raw(self, context: str, attempt_index: int = 0) -> str: | |
| """Generate sekali. Return raw string dari model.""" | |
| inputs = self.tokenizer.encode( | |
| context, | |
| return_tensors='pt', | |
| truncation=True, | |
| max_length=self.max_context_tokens, | |
| ).to(self.device) | |
| input_length = inputs.shape[1] | |
| model_max_len = getattr(self.model.config, 'max_position_embeddings', 1024) | |
| max_new_tokens = min( | |
| self.gen_config["max_new_tokens"], | |
| model_max_len - input_length - 10, | |
| ) | |
| if max_new_tokens <= 0: | |
| keep_tokens = model_max_len - 50 | |
| inputs = inputs[:, -keep_tokens:] | |
| input_length = inputs.shape[1] | |
| max_new_tokens = 30 | |
| # Temperature bervariasi antar attempt → diversity | |
| base_temp = self.gen_config["temperature"] | |
| temp_var = base_temp + (attempt_index - 1) * 0.05 | |
| temp_var = max(0.5, min(temp_var, 1.1)) | |
| kwargs = { | |
| "max_new_tokens": max(1, max_new_tokens), | |
| "temperature": temp_var, | |
| "top_p": self.gen_config["top_p"], | |
| "top_k": self.gen_config["top_k"], | |
| "do_sample": True, | |
| "pad_token_id": self.tokenizer.eos_token_id, | |
| "eos_token_id": self.tokenizer.eos_token_id, | |
| "repetition_penalty": self.gen_config["repetition_penalty"], | |
| "no_repeat_ngram_size": self.gen_config["no_repeat_ngram_size"], | |
| "use_cache": True, | |
| } | |
| with torch.no_grad(): | |
| outputs = self.model.generate(inputs, **kwargs) | |
| generated_tokens = outputs[0, input_length:] | |
| return self.tokenizer.decode(generated_tokens, skip_special_tokens=True) | |
| # ================================================================== | |
| # Response scoring | |
| # ================================================================== | |
| def _score_response(self, response: str, user_input: str) -> float: | |
| """ | |
| Skor 0.0 - 1.0. Semakin tinggi semakin baik. | |
| Penalty untuk: junk phrases, repetitive, boilerplate, terlalu pendek. | |
| """ | |
| if not response or not response.strip(): | |
| return 0.0 | |
| score = 1.0 | |
| words = response.lower().split() | |
| word_count = len(words) | |
| # 1. Panjang wajar: 3 - 40 kata ideal | |
| if word_count < 3: | |
| score -= 0.5 | |
| elif word_count < 5: | |
| score -= 0.2 | |
| elif word_count > 40: | |
| score -= 0.3 | |
| elif word_count > 60: | |
| score -= 0.6 | |
| # 2. Uniqueness: kalau banyak kata diulang = jelek | |
| if word_count >= 5: | |
| unique_ratio = len(set(words)) / word_count | |
| if unique_ratio < 0.5: | |
| score -= 0.5 | |
| elif unique_ratio < 0.65: | |
| score -= 0.2 | |
| # 3. Junk phrases | |
| lower = response.lower() | |
| junk_count = sum(1 for p in _JUNK_PHRASES if p in lower) | |
| score -= junk_count * 0.25 | |
| # 4. Boilerplate | |
| boilerplate_count = sum(1 for m in _BOILERPLATE_MARKERS if m in lower) | |
| score -= boilerplate_count * 0.15 | |
| # 5. Terlalu banyak tanda tanya | |
| q_count = response.count("?") | |
| if q_count >= 3: | |
| score -= 0.3 | |
| elif q_count >= 5: | |
| score -= 0.6 | |
| # 6. Terlalu banyak koma = kalimat tidak terstruktur | |
| comma_count = response.count(",") | |
| if word_count > 0 and comma_count / max(word_count, 1) > 0.15: | |
| score -= 0.15 | |
| # 7. Bonus: ada kata kunci user di response (relevan) | |
| user_words = set(w.lower() for w in re.findall(r'\w+', user_input) if len(w) > 3) | |
| if user_words: | |
| overlap = len(user_words & set(words)) | |
| if overlap > 0: | |
| score += min(0.2, overlap * 0.05) | |
| # 8. Bonus: respon pendek 5-25 kata = ideal | |
| if 5 <= word_count <= 25 and not any(p in lower for p in _JUNK_PHRASES): | |
| score += 0.1 | |
| return max(0.0, min(1.0, score)) | |
| # ================================================================== | |
| # Response cleaning (agresif) | |
| # ================================================================== | |
| def _clean_response(self, response: str, user_input: str) -> str: | |
| if not response: | |
| return "" | |
| # 1. Buang user_input yang ikut ke-generate | |
| if user_input and len(user_input) > 3: | |
| response = response.replace(user_input, '').strip() | |
| # 2. Buang markup sisa | |
| response = re.sub(r'^(AI:|User:)', '', response).strip() | |
| for marker in ['User:', 'AI:', '<|endoftext|>', '<|end|>']: | |
| pos = response.find(marker) | |
| if pos != -1: | |
| response = response[:pos].strip() | |
| response = ( | |
| response.replace('<|endoftext|>', '') | |
| .replace('<|end|>', '') | |
| .strip() | |
| ) | |
| # 3. Rapikan whitespace | |
| response = ' '.join(response.split()) | |
| # 4. Fix typo umum yang konsisten muncul | |
| typo_fixes = { | |
| "Seminta": "Semoga", | |
| "Semohon": "Semoga", | |
| "semiap": "siap", | |
| } | |
| for typo, fix in typo_fixes.items(): | |
| response = response.replace(typo, fix) | |
| # 5. Cutoff di junk phrase pertama | |
| lower = response.lower() | |
| cutoff = len(response) | |
| for phrase in _JUNK_PHRASES: | |
| pos = lower.find(phrase) | |
| if pos != -1 and pos < cutoff: | |
| cutoff = pos | |
| if cutoff < len(response): | |
| response = response[:cutoff].rstrip(" ,;.") | |
| # 6. Cutoff di boilerplate pertama (setelah 25% awal) | |
| lower = response.lower() | |
| cutoff = len(response) | |
| min_pos = int(len(response) * 0.25) | |
| for marker in _BOILERPLATE_MARKERS: | |
| pos = lower.find(marker, min_pos) | |
| if pos != -1 and pos < cutoff: | |
| cutoff = pos | |
| if cutoff < len(response): | |
| response = response[:cutoff].rstrip(" ,;.") | |
| # 7. Fix punctuation aneh | |
| response = re.sub(r'[;:,]+\.', '.', response) | |
| response = re.sub(r'\.\.+', '.', response) | |
| response = re.sub(r'[,]{2,}', ',', response) | |
| response = re.sub(r'[;]{2,}', ';', response) | |
| response = re.sub(r'\s+([.,!?;:])', r'\1', response) | |
| response = re.sub(r'([.,!?;:])([^\s\d])', r'\1 \2', response) | |
| response = re.sub(r'\s+\.', '.', response) | |
| # 8. Batas kalimat: max 2 | |
| sentences = re.split(r'(?<=[.!?])\s+', response) | |
| if len(sentences) > 2: | |
| response = ' '.join(sentences[:2]).strip() | |
| if response and response[-1] not in ".!?": | |
| response += "." | |
| # 9. Hard cap karakter | |
| if len(response) > 250: | |
| response = response[:250].rsplit(' ', 1)[0] + "..." | |
| return response.strip() | |
| # ================================================================== | |
| # Streaming (tetap ada) | |
| # ================================================================== | |
| def _generate_streaming(self, user_input: str) -> str: | |
| if not self.model or not self.tokenizer: | |
| return "Model belum siap" | |
| try: | |
| context = self._build_context(user_input) | |
| inputs = self.tokenizer.encode( | |
| context, return_tensors='pt', | |
| truncation=True, max_length=self.max_context_tokens, | |
| ).to(self.device) | |
| input_length = inputs.shape[1] | |
| model_max_len = getattr(self.model.config, 'max_position_embeddings', 1024) | |
| max_new_tokens = min( | |
| self.gen_config["max_new_tokens"], | |
| model_max_len - input_length - 10, | |
| ) | |
| if max_new_tokens <= 0: | |
| max_new_tokens = 30 | |
| streamer = TextIteratorStreamer( | |
| self.tokenizer, skip_prompt=True, skip_special_tokens=True | |
| ) | |
| gen_kwargs = { | |
| "input_ids": inputs, | |
| "max_new_tokens": max_new_tokens, | |
| "temperature": self.gen_config["temperature"], | |
| "top_p": self.gen_config["top_p"], | |
| "top_k": self.gen_config["top_k"], | |
| "repetition_penalty": self.gen_config["repetition_penalty"], | |
| "no_repeat_ngram_size": self.gen_config["no_repeat_ngram_size"], | |
| "do_sample": True, | |
| "pad_token_id": self.tokenizer.eos_token_id, | |
| "eos_token_id": self.tokenizer.eos_token_id, | |
| "streamer": streamer, | |
| } | |
| thread = Thread(target=self.model.generate, kwargs=gen_kwargs) | |
| thread.start() | |
| console.print("[magenta]AI[/]: ", end="") | |
| chunks: List[str] = [] | |
| for new_text in streamer: | |
| console.print(rich_escape(new_text), end="") | |
| chunks.append(new_text) | |
| console.print() | |
| console.print() | |
| response = "".join(chunks) | |
| response = self._clean_response(response, user_input) | |
| return response if response else _FALLBACKS["unknown"] | |
| except torch.cuda.OutOfMemoryError: | |
| console.print(Theme.error(" GPU OOM. Fallback ke CPU...")) | |
| try: | |
| self.model = self.model.to('cpu') | |
| torch.cuda.empty_cache() | |
| self.device = 'cpu' | |
| return self._generate_best_of_n(user_input) | |
| except Exception as e: | |
| return f"Error: {str(e)[:100]}" | |
| except Exception as e: | |
| error_logger.error(f"Streaming error: {e}") | |
| return f"Error: {str(e)[:120]}" | |
| # ================================================================== | |
| # Model loading (tidak diubah) | |
| # ================================================================== | |
| def _load_model(self, model_path: str) -> None: | |
| is_peft = False | |
| HAS_PEFT = False | |
| try: | |
| from peft import PeftModel | |
| HAS_PEFT = True | |
| except ImportError: | |
| pass | |
| if os.path.exists(os.path.join(model_path, "adapter_config.json")): | |
| is_peft = True | |
| base_info_path = os.path.join(model_path, "base_model_info.json") | |
| if os.path.exists(base_info_path): | |
| try: | |
| with open(base_info_path, 'r') as f: | |
| base_info = json.load(f) | |
| if base_info.get('use_peft', False): | |
| is_peft = True | |
| except Exception: | |
| pass | |
| if is_peft and HAS_PEFT: | |
| self._load_peft_model(model_path) | |
| else: | |
| self._load_full_model(model_path) | |
| def _load_peft_model(self, model_path: str) -> None: | |
| try: | |
| from peft import PeftModel | |
| base_info_path = os.path.join(model_path, "base_model_info.json") | |
| if os.path.exists(base_info_path): | |
| with open(base_info_path, 'r') as f: | |
| base_info = json.load(f) | |
| base_model_name = base_info.get('base_model', 'gpt2') | |
| else: | |
| base_model_name = 'gpt2' | |
| self.tokenizer = AutoTokenizer.from_pretrained(model_path) | |
| if self.tokenizer.pad_token is None: | |
| self.tokenizer.pad_token = self.tokenizer.eos_token | |
| torch_dtype = torch.float16 if self.device == 'cuda' else torch.float32 | |
| self.model = AutoModelForCausalLM.from_pretrained( | |
| base_model_name, | |
| torch_dtype=torch_dtype, | |
| low_cpu_mem_usage=True, | |
| ) | |
| self.model = PeftModel.from_pretrained(self.model, model_path) | |
| self.is_peft_model = True | |
| self.model = self.model.to(self.device if self.device != 'cpu' else 'cpu') | |
| console.print(Theme.success(" PEFT model + adapter loaded!")) | |
| except Exception as e: | |
| console.print(Theme.warning(f"PEFT load failed: {e}. Fallback ke full model.")) | |
| self._load_full_model(model_path) | |
| def _load_full_model(self, model_path: str) -> None: | |
| torch_dtype = torch.float16 if self.device == 'cuda' else torch.float32 | |
| self.tokenizer = AutoTokenizer.from_pretrained(model_path) | |
| if self.tokenizer.pad_token is None: | |
| self.tokenizer.pad_token = self.tokenizer.eos_token | |
| self.model = AutoModelForCausalLM.from_pretrained( | |
| model_path, | |
| torch_dtype=torch_dtype, | |
| low_cpu_mem_usage=True, | |
| ) | |
| self.model = self.model.to(self.device if self.device != 'cpu' else 'cpu') | |
| self.is_peft_model = False | |
| console.print(Theme.success(" Full model loaded!")) | |
| def _sync_gen_config_from_manager(self) -> None: | |
| try: | |
| cfg = self.model_manager.get_generation_config() | |
| self.max_context_tokens = cfg.get('max_context_length', 768) | |
| # Hanya temperature/top_p yang di-sync; parameter "aman" tetap | |
| # pakai default chat.py (max_new_tokens kecil, rep_penalty tinggi). | |
| if "temperature" in cfg: | |
| self.gen_config["temperature"] = min(cfg["temperature"], 0.9) | |
| if "top_p" in cfg: | |
| self.gen_config["top_p"] = cfg["top_p"] | |
| except Exception as e: | |
| debug_logger.debug(f"Sync gen config: {e}") | |
| # ================================================================== | |
| # Context building | |
| # ================================================================== | |
| def _build_context(self, user_input: str) -> str: | |
| """Context ketat: system + max 2 turn terakhir + user input.""" | |
| context = f"{self.system_prompt}\n\n" | |
| total_tokens = len(self.tokenizer.encode(context)) | |
| for h in list(self.conversation_history)[-2:]: | |
| turn = f"User: {h['user'][:60]}\nAI: {h['ai'][:120]}\n\n" | |
| turn_tokens = len(self.tokenizer.encode(turn)) | |
| if total_tokens + turn_tokens > self.max_context_tokens - 100: | |
| break | |
| context += turn | |
| total_tokens += turn_tokens | |
| context += f"User: {user_input}\nAI:" | |
| return context | |
| # ================================================================== | |
| # Stats | |
| # ================================================================== | |
| def _update_stats(self, elapsed_ms: float, response: str) -> None: | |
| self.performance_stats["total_generations"] += 1 | |
| n = self.performance_stats["total_generations"] | |
| avg = self.performance_stats["avg_time_ms"] | |
| self.performance_stats["avg_time_ms"] = (avg * (n - 1) + elapsed_ms) / n | |
| self.performance_stats["max_response_time"] = max( | |
| self.performance_stats["max_response_time"], elapsed_ms | |
| ) | |
| self.performance_stats["min_response_time"] = min( | |
| self.performance_stats["min_response_time"], elapsed_ms | |
| ) | |
| self.performance_stats["total_tokens_generated"] += len(response.split()) | |
| def _show_stats(self) -> None: | |
| console.print("\n" + Theme.header(" Chat Performance Statistics:")) | |
| table = Table(title="Performance Stats", box=box.ROUNDED) | |
| table.add_column("Metric", style="cyan") | |
| table.add_column("Value", style="green") | |
| s = self.performance_stats | |
| table.add_row("Total Generations", str(s["total_generations"])) | |
| table.add_row("Avg Response Time", f"{s['avg_time_ms']:.1f}ms") | |
| table.add_row("Max Response Time", f"{s['max_response_time']:.1f}ms") | |
| min_t = s["min_response_time"] | |
| if min_t == float('inf'): | |
| min_t = 0.0 | |
| table.add_row("Min Response Time", f"{min_t:.1f}ms") | |
| table.add_row("Total Tokens Generated", f"{s['total_tokens_generated']:,}") | |
| table.add_row("Total Tokens Processed", f"{s['total_tokens_processed']:,}") | |
| table.add_row("Fallbacks Triggered", str(s.get("fallbacks_triggered", 0))) | |
| table.add_row("Regenerations", str(s.get("regenerations", 0))) | |
| table.add_row("Device", self.device.upper()) | |
| table.add_row("PEFT Model", "Yes" if self.is_peft_model else "No") | |
| table.add_row("History Length", str(len(self.conversation_history))) | |
| table.add_row("Streaming", "ON" if self.gen_config["streaming"] else "OFF") | |
| console.print(table) | |
| # Config | |
| table2 = Table(title="Generation Config", box=box.ROUNDED) | |
| table2.add_column("Parameter", style="cyan") | |
| table2.add_column("Value", style="green") | |
| for k, v in self.gen_config.items(): | |
| table2.add_row(k, str(v)) | |
| console.print(table2) | |
| # ================================================================== | |
| # History | |
| # ================================================================== | |
| def _show_history(self) -> None: | |
| if not self.conversation_history: | |
| console.print(Theme.warning("Belum ada riwayat percakapan")) | |
| return | |
| console.print("\n" + Theme.header(" Conversation History:")) | |
| for idx, item in enumerate(self.conversation_history, 1): | |
| ts = datetime.fromtimestamp(item.get("timestamp", 0)).strftime("%H:%M") | |
| console.print(f"\n[green][{idx}] User:[/green] [dim]({ts})[/dim]") | |
| console.print(f" {rich_escape(str(item['user']))}") | |
| console.print(f"[magenta] AI:[/magenta]") | |
| console.print(f" {rich_escape(str(item['ai']))}") | |
| # ================================================================== | |
| # Session save/load | |
| # ================================================================== | |
| def _save_session(self) -> None: | |
| if not self.conversation_history: | |
| console.print(Theme.warning("Belum ada percakapan untuk disimpan")) | |
| return | |
| ts = datetime.now().strftime("%Y%m%d_%H%M%S") | |
| out = _SESSION_DIR / f"session_{ts}.json" | |
| try: | |
| data = { | |
| "saved_at": datetime.now().isoformat(), | |
| "persona": self.system_prompt, | |
| "gen_config": self.gen_config, | |
| "stats": self.performance_stats, | |
| "history": [ | |
| { | |
| "user": h["user"], | |
| "ai": h["ai"], | |
| "timestamp": h.get("timestamp", 0), | |
| } | |
| for h in self.conversation_history | |
| ], | |
| } | |
| with open(out, "w", encoding="utf-8") as f: | |
| json.dump(data, f, indent=2, ensure_ascii=False) | |
| console.print(Theme.success(f"✓ Session disimpan: {out.name}")) | |
| except Exception as e: | |
| console.print(Theme.error(f"Gagal save: {e}")) | |
| def _list_sessions(self) -> None: | |
| sessions = sorted(_SESSION_DIR.glob("session_*.json"), reverse=True) | |
| if not sessions: | |
| console.print(Theme.warning("Belum ada session tersimpan")) | |
| return | |
| table = Table(title="Saved Sessions", box=box.ROUNDED) | |
| table.add_column("No", style="cyan") | |
| table.add_column("Nama", style="green") | |
| table.add_column("Turns", style="blue") | |
| table.add_column("Waktu", style="dim") | |
| for i, s in enumerate(sessions, 1): | |
| try: | |
| with open(s) as f: | |
| data = json.load(f) | |
| turns = len(data.get("history", [])) | |
| ts = data.get("saved_at", "")[:19].replace("T", " ") | |
| table.add_row(str(i), s.name, str(turns), ts) | |
| except Exception: | |
| table.add_row(str(i), s.name, "?", "?") | |
| console.print(table) | |
| console.print(Theme.dim("Load dengan: load <nomor>")) | |
| def _load_session(self, idx: int) -> None: | |
| sessions = sorted(_SESSION_DIR.glob("session_*.json"), reverse=True) | |
| if not sessions: | |
| console.print(Theme.warning("Belum ada session")) | |
| return | |
| if idx < 1 or idx > len(sessions): | |
| console.print(Theme.error("Nomor tidak valid")) | |
| return | |
| path = sessions[idx - 1] | |
| try: | |
| with open(path, encoding="utf-8") as f: | |
| data = json.load(f) | |
| self.conversation_history.clear() | |
| for h in data.get("history", []): | |
| self.conversation_history.append({ | |
| "user": h["user"], | |
| "ai": h["ai"], | |
| "timestamp": h.get("timestamp", time.time()), | |
| }) | |
| if data.get("persona"): | |
| self.system_prompt = data["persona"] | |
| console.print(Theme.success( | |
| f"✓ Session di-load: {len(self.conversation_history)} turns" | |
| )) | |
| except Exception as e: | |
| console.print(Theme.error(f"Gagal load: {e}")) | |
| # ================================================================== | |
| # Export chat ke markdown | |
| # ================================================================== | |
| def _export_chat(self) -> None: | |
| if not self.conversation_history: | |
| console.print(Theme.warning("Belum ada percakapan")) | |
| return | |
| ts = datetime.now().strftime("%Y%m%d_%H%M%S") | |
| out = Path(f"/content/data/chat_export_{ts}.md") | |
| try: | |
| lines = [ | |
| "# Chat Export", | |
| "", | |
| f"**Tanggal:** {datetime.now().strftime('%Y-%m-%d %H:%M:%S')}", | |
| f"**Model:** {self.model_manager.get_current_model()}", | |
| f"**Device:** {self.device.upper()}", | |
| f"**Total turns:** {len(self.conversation_history)}", | |
| "", | |
| "---", | |
| "", | |
| ] | |
| for h in self.conversation_history: | |
| ts_msg = datetime.fromtimestamp(h.get("timestamp", 0)).strftime("%H:%M:%S") | |
| lines.append(f"### User ({ts_msg})") | |
| lines.append("") | |
| lines.append(h["user"]) | |
| lines.append("") | |
| lines.append(f"### AI") | |
| lines.append("") | |
| lines.append(h["ai"]) | |
| lines.append("") | |
| lines.append("---") | |
| lines.append("") | |
| out.write_text("\n".join(lines), encoding="utf-8") | |
| console.print(Theme.success(f"✓ Chat di-export: {out}")) | |
| except Exception as e: | |
| console.print(Theme.error(f"Gagal export: {e}")) | |
| # ================================================================== | |
| # Help | |
| # ================================================================== | |
| def _show_help(self) -> None: | |
| help_text = """ | |
| [cyan]KONVERSASI:[/cyan] | |
| (ketik langsung) - kirim pesan ke AI | |
| [cyan]RIWAYAT:[/cyan] | |
| history / hist - tampilkan riwayat | |
| undo / u - hapus turn terakhir | |
| regen / rg - jawab ulang pesan terakhir | |
| clear / xlear - hapus semua riwayat | |
| [cyan]PARAMETER (runtime):[/cyan] | |
| temp 0.9 - ubah temperature (0.01-2.0) | |
| topp 0.85 - ubah top-p (0.0-1.0) | |
| maxtok 60 - ubah max tokens (10-500) | |
| attempts 3 - ubah Best-of-N (1-10) | |
| stream on/off - aktifkan streaming output | |
| [cyan]PERSONA:[/cyan] | |
| persona <teks> - ganti system prompt | |
| persona reset - kembali ke default | |
| persona - lihat persona sekarang | |
| [cyan]SESSION:[/cyan] | |
| save / sessions / load N / export | |
| [cyan]LAIN-LAIN:[/cyan] | |
| stats / s - statistik performa | |
| help / hlp - tampilkan bantuan ini | |
| exit / q - keluar | |
| [cyan]TIPS:[/cyan] | |
| - Greeting/thanks/bye dijawab langsung tanpa model (cepat, akurat) | |
| - Best-of-N: model generate 3 jawaban, dipilih yang terbaik | |
| - Kalau jawaban jelek, sistem otomatis fallback ke template sopan | |
| - Gunakan 'regen' kalau ingin jawaban berbeda | |
| - 'attempts 5' bikin lebih variatif (tapi lebih lambat) | |
| """ | |
| console.print(Panel(help_text.strip(), title="HELP", style="yellow")) |