| import os |
| import torch |
| import json |
| from datetime import datetime |
|
|
| from config import * |
| from model import PegeModel |
| from tokenizer import TurkishTokenizer |
|
|
| class PegeChat: |
| def __init__(self): |
| self.tokenizer = TurkishTokenizer(vocab_size=VOCAB_SIZE) |
| self.model = None |
| self.feedback_buffer = [] |
| self.conversation_history = [] |
|
|
| self.load_model() |
|
|
| def load_model(self): |
| """Model ve tokenizer yükle""" |
| |
| tokenizer_path = os.path.expanduser(f"{MODEL_DIR}/tokenizer.pkl") |
| if os.path.exists(tokenizer_path): |
| self.tokenizer.load(tokenizer_path) |
| else: |
| print("UYARI: Tokenizer eğitilmemiş! Önce python train.py çalıştır.") |
| return False |
|
|
| vocab_size = len(self.tokenizer.stoi) |
|
|
| |
| self.model = PegeModel(vocab_size).to(device) |
|
|
| |
| checkpoint_files = glob.glob(os.path.expanduser(f"{CHECKPOINT_DIR}/checkpoint_*.pt")) |
| model_path = os.path.expanduser(f"{MODEL_DIR}/pege_final.pt") |
| finetuned_path = os.path.expanduser(f"{MODEL_DIR}/pege_finetuned.pt") |
|
|
| if os.path.exists(finetuned_path): |
| checkpoint = torch.load(finetuned_path, map_location=device) |
| self.model.load_state_dict(checkpoint['model_state']) |
| print(f"Model loaded from: {finetuned_path}") |
| elif checkpoint_files: |
| latest = max(checkpoint_files, key=os.path.getctime) |
| checkpoint = torch.load(latest, map_location=device) |
| self.model.load_state_dict(checkpoint['model_state']) |
| print(f"Model loaded from: {latest}") |
| elif os.path.exists(model_path): |
| checkpoint = torch.load(model_path, map_location=device) |
| self.model.load_state_dict(checkpoint['model_state']) |
| print(f"Model loaded from: {model_path}") |
| else: |
| print("UYARI: Eğitilmiş model bulunamadı!") |
| print("Önce python train.py çalıştır.") |
| return False |
|
|
| self.model.eval() |
| return True |
|
|
| def generate_response(self, user_input, temperature=0.8): |
| """Kullanıcı input'una cevap üret""" |
| |
| prompt = f"{self.tokenizer.USER}{user_input}{self.tokenizer.ASSISTANT}" |
|
|
| |
| context = "" |
| for turn in self.conversation_history[-3:]: |
| context += f"{self.tokenizer.USER}{turn['input']}{self.tokenizer.ASSISTANT}{turn['output']}" |
|
|
| full_prompt = context + prompt |
|
|
| input_ids = self.tokenizer.encode(full_prompt) |
| input_tensor = torch.tensor([input_ids], dtype=torch.long).to(device) |
|
|
| stop_ids = [ |
| self.tokenizer.stoi.get(self.tokenizer.USER, -1), |
| self.tokenizer.stoi.get(self.tokenizer.EOS, -1), |
| ] |
| output_ids = self.model.generate( |
| input_tensor, |
| max_new_tokens=50, |
| temperature=temperature, |
| top_k=40, |
| repetition_penalty=1.5, |
| stop_tokens=stop_ids, |
| ) |
|
|
| |
| new_tokens = output_ids[0][len(input_ids):].tolist() |
| response = self.tokenizer.decode(new_tokens) |
|
|
| |
| response = response.replace(self.tokenizer.ASSISTANT, "").strip() |
|
|
| return response |
|
|
| def save_feedback(self, user_input, response, feedback): |
| """Geri bildirimi kaydet""" |
| entry = { |
| 'timestamp': datetime.now().isoformat(), |
| 'input': user_input, |
| 'output': response, |
| 'feedback': feedback |
| } |
|
|
| |
| self.feedback_buffer.append(entry) |
|
|
| |
| feedback_path = os.path.expanduser(FEEDBACK_FILE) |
| os.makedirs(os.path.dirname(feedback_path), exist_ok=True) |
|
|
| with open(feedback_path, 'a', encoding='utf-8') as f: |
| f.write(json.dumps(entry, ensure_ascii=False) + '\n') |
|
|
| print(f"✓ Feedback kaydedildi: {feedback}") |
|
|
| def run(self): |
| """Ana döngü""" |
| if not self.load_model(): |
| return |
|
|
| print("\n" + "=" * 50) |
| print("PEGE CHAT") |
| print("=" * 50) |
| print("Komutlar:") |
| print(" /good - Cevabı beğendin, öğren") |
| print(" /bad - Cevabı beğenmedin, tekrar dene") |
| print(" /retry - Yeni cevap iste") |
| print(" /clear - Konuşma geçmişini temizle") |
| print(" /quit - Çıkış") |
| print("=" * 50 + "\n") |
|
|
| last_input = None |
| last_response = None |
|
|
| while True: |
| try: |
| user_input = input("\nSen: ").strip() |
|
|
| if not user_input: |
| continue |
|
|
| |
| if user_input == "/quit": |
| print("Görüşmek üzere!") |
| break |
|
|
| if user_input == "/clear": |
| self.conversation_history = [] |
| print("Konuşma geçmişi temizlendi.") |
| continue |
|
|
| if user_input == "/retry": |
| if last_input: |
| print("Yeni cevap deneniyor...") |
| response = self.generate_response(last_input, temperature=0.9) |
| print(f"\nPege: {response}") |
| last_response = response |
| continue |
|
|
| if user_input == "/good": |
| if last_input and last_response: |
| self.save_feedback(last_input, last_response, "good") |
| self.conversation_history.append({ |
| 'input': last_input, |
| 'output': last_response |
| }) |
| print("Öğrenildi! Bu tarz cevapları daha çok vereceğim.") |
| continue |
|
|
| if user_input == "/bad": |
| if last_input and last_response: |
| self.save_feedback(last_input, last_response, "bad") |
| print("Noted. Bu cevabı beğenmedin.") |
| print("Yeni cevap deneniyor...") |
|
|
| |
| response = self.generate_response(last_input, temperature=1.1) |
| print(f"\nPege: {response}") |
| last_response = response |
| continue |
|
|
| |
| last_input = user_input |
| print("Pege düşünüyor...", end=" ") |
|
|
| response = self.generate_response(user_input, temperature=0.5) |
| last_response = response |
|
|
| print(f"\nPege: {response}") |
| print("\n(good/bad/retry/quit)") |
|
|
| except KeyboardInterrupt: |
| print("\nÇıkılıyor...") |
| break |
| except Exception as e: |
| print(f"\nHata: {e}") |
|
|
| import glob |
|
|
| if __name__ == "__main__": |
| chat = PegeChat() |
| chat.run() |