import torch from transformers import AutoModelForCausalLM, AutoTokenizer, StoppingCriteria, StoppingCriteriaList # Settings MODEL_PATH = "DrontChat-200m" # path to local model DEVICE = "cuda" if torch.cuda.is_available() else "cpu" TEMPERATURE = 0.3 MAX_NEW_TOKENS = 256 TOP_P = 0.90 TOP_K = 50 class StopOnTokens(StoppingCriteria): """Class to stop generation when encountering stop tokens""" def __init__(self, stop_token_ids): self.stop_token_ids = set(stop_token_ids) def __call__(self, input_ids, scores, **kwargs): # Check the last generated token if input_ids.shape[-1] > 0: last_token = input_ids[0, -1].item() if last_token in self.stop_token_ids: return True return False class LocalChatBot: def __init__(self, model_path): print(f"Loading model from {model_path}...") # Load tokenizer and model self.tokenizer = AutoTokenizer.from_pretrained( model_path, trust_remote_code=True, padding_side="left" ) # Define special tokens special_tokens = { "pad_token": "<|endoftext|>", "eos_token": "<|endoftext|>", "sep_token": "<|endoftext|>", "additional_special_tokens": ["<|user|>", "<|assistant|>", "", ""] } # Add special tokens self.tokenizer.add_special_tokens(special_tokens) self.model = AutoModelForCausalLM.from_pretrained( model_path, torch_dtype=torch.float16 if DEVICE == "cuda" else torch.float32, device_map="auto", trust_remote_code=True, low_cpu_mem_usage=True ) # If new tokens were added, resize embeddings if len(self.tokenizer) > self.model.config.vocab_size: self.model.resize_token_embeddings(len(self.tokenizer)) self.model.eval() # Get stop token IDs self.stop_token_ids = self._get_stop_token_ids() if DEVICE == "cuda": print(f"Model loaded on GPU: {torch.cuda.get_device_name()}") else: print("WARNING: CUDA not available, using CPU!") def _get_stop_token_ids(self): """Get IDs of all stop tokens""" stop_tokens = [ "<|endoftext|>", "<|user|>", # Stop if model starts generating for user "<|assistant|>", # Stop if model generates new response "", # Stop if model starts new system prompt ] stop_ids = [] for token in stop_tokens: token_id = self.tokenizer.convert_tokens_to_ids(token) if token_id is not None and token_id != -1: stop_ids.append(token_id) print(f"Stop token '{token}' -> ID: {token_id}") else: print(f"Warning: token '{token}' not found in tokenizer") return stop_ids def format_prompt(self, system_message, user_input, history=[]): """Format prompt with conversation history""" prompt = f"{system_message}" # Add conversation history for user_msg, assistant_msg in history: prompt += f"<|user|>{user_msg}<|endoftext|>" prompt += f"<|assistant|>{assistant_msg}<|endoftext|>" # Add current message prompt += f"<|user|>{user_input}<|endoftext|>" prompt += "<|assistant|>" return prompt def generate_response(self, prompt): """Generate model response with proper stopping""" inputs = self.tokenizer.encode( prompt, return_tensors="pt", truncation=True, max_length=2048, add_special_tokens=False ).to(DEVICE) # Create stopping criteria stopping_criteria = StoppingCriteriaList([StopOnTokens(self.stop_token_ids)]) with torch.no_grad(): outputs = self.model.generate( inputs, max_new_tokens=MAX_NEW_TOKENS, temperature=TEMPERATURE, do_sample=True if TEMPERATURE > 0 else False, top_p=TOP_P, top_k=TOP_K, pad_token_id=self.tokenizer.pad_token_id, eos_token_id=self.tokenizer.eos_token_id, repetition_penalty=1.1, num_return_sequences=1, stopping_criteria=stopping_criteria, # Add stopping criteria ) # Decode only new tokens response = self.tokenizer.decode( outputs[0][inputs.shape[1]:], skip_special_tokens=True ).strip() # Additional cleanup from possible markers response = self._clean_response(response) return response def _clean_response(self, response): """Clean response from service tokens""" # List of markers for cleanup markers = [ "<|endoftext|>", "<|user|>", "<|assistant|>", "", "" ] for marker in markers: if marker in response: response = response.split(marker)[0].strip() return response def chat(self): """Interactive chat""" print("\n" + "=" * 50) print("Local chat bot started!") print(f"Temperature: {TEMPERATURE}") print(f"Device: {DEVICE}") print("Commands: 'clear' - clear history, 'exit' - exit") print("=" * 50 + "\n") system_message = "You are a AI, you can smol talk, you have name DrontAI." history = [] while True: try: user_input = input("You: ").strip() if not user_input: continue if user_input.lower() == 'exit': print("Goodbye!") break if user_input.lower() == 'clear': history = [] print("Conversation history cleared.") continue if user_input.lower().startswith('system:'): system_message = user_input[7:].strip() print(f"System message updated: {system_message}") continue # Format prompt prompt = self.format_prompt(system_message, user_input, history) # Generate response response = self.generate_response(prompt) # Check if response is not empty if not response: response = "(empty response)" # Save to history history.append((user_input, response)) # Limit history to last 5 exchanges if len(history) > 5: history = history[-5:] print(f"Bot: {response}\n") except KeyboardInterrupt: print("\nInterrupted by user.") break except Exception as e: print(f"Error: {e}") continue def main(): """Main function""" try: # Check CUDA availability if torch.cuda.is_available(): print(f"CUDA available: {torch.cuda.get_device_name(0)}") print(f"GPU Memory: {torch.cuda.get_device_properties(0).total_memory / 1e9:.1f} GB") # Clear CUDA cache torch.cuda.empty_cache() else: print("CUDA not available. Will use CPU (slow).") # Create and run bot bot = LocalChatBot(MODEL_PATH) bot.chat() except Exception as e: print(f"Critical error: {e}") import traceback traceback.print_exc() if __name__ == "__main__": main()