# use.py - Chat with VDrontMoE model (text continuation) import torch import argparse import torch.nn.functional as F import os from transformers import AutoTokenizer from create import VDrontMoEConfig, VDrontMoEModel # --- Configuration --- MODEL_DIR = "./VDrontMoE-2m-5e" DEVICE = "cuda" if torch.cuda.is_available() else "cpu" TEMPERATURE = 0.4 MAX_NEW_TOKENS = 100 TOP_P = 0.9 TOP_K = 50 REPETITION_PENALTY = 1.1 # --- Load model --- print(f"Loading model from {MODEL_DIR}...") tokenizer = AutoTokenizer.from_pretrained(MODEL_DIR) if tokenizer.pad_token is None: tokenizer.pad_token = tokenizer.eos_token model = VDrontMoEModel.from_pretrained(MODEL_DIR).to(DEVICE) model.eval() print(f"Model loaded! Temperature: {TEMPERATURE}") def generate(prompt: str, max_tokens: int = MAX_NEW_TOKENS) -> str: """Continue text from prompt.""" # Tokenize prompt input_ids = tokenizer.encode(prompt, return_tensors="pt").to(DEVICE) generated = input_ids.clone() with torch.no_grad(): for _ in range(max_tokens): # Truncate if too long if generated.shape[1] > 512: generated = generated[:, -512:] # Forward pass outputs = model(generated) logits = outputs["logits"][:, -1, :] # Take last token # Temperature logits = logits / TEMPERATURE # Top-K filter if TOP_K > 0: top_k_values, _ = torch.topk(logits, min(TOP_K, logits.size(-1))) logits[logits < top_k_values[:, -1:]] = float('-inf') # Top-P (nucleus) filter if TOP_P < 1.0: sorted_logits, sorted_indices = torch.sort(logits, descending=True) cumulative_probs = torch.cumsum(F.softmax(sorted_logits, dim=-1), dim=-1) # Remove tokens with cumulative prob > TOP_P sorted_indices_to_remove = cumulative_probs > TOP_P sorted_indices_to_remove[:, 1:] = sorted_indices_to_remove[:, :-1].clone() sorted_indices_to_remove[:, 0] = False indices_to_remove = sorted_indices_to_remove.scatter(1, sorted_indices, sorted_indices_to_remove) logits[indices_to_remove] = float('-inf') # Repetition penalty if REPETITION_PENALTY != 1.0: for token_id in set(generated[0].tolist()[-10:]): # Last 10 tokens if logits[0, token_id] > 0: logits[0, token_id] /= REPETITION_PENALTY else: logits[0, token_id] *= REPETITION_PENALTY # Sample next token probs = F.softmax(logits, dim=-1) next_token = torch.multinomial(probs, num_samples=1) # Add to generated generated = torch.cat([generated, next_token], dim=-1) # Check for end of text if next_token.item() == tokenizer.eos_token_id: break # Decode full_text = tokenizer.decode(generated[0], skip_special_tokens=True) return full_text def chat(): """Interactive text continuation mode.""" print("\n" + "=" * 60) print("VDrontMoE-2m-5e - Text Continuation Mode") print("=" * 60) print("Commands:") print(" /temp - Set temperature (0.1 - 2.0)") print(" /max - Set max new tokens (10 - 500)") print(" /clear - Clear screen") print(" /exit - Exit") print("=" * 60) global TEMPERATURE, MAX_NEW_TOKENS while True: try: # Get prompt prompt = input("\nPrompt: ").strip() # Commands if prompt.startswith("/"): parts = prompt.split() cmd = parts[0] if cmd == "/exit": print("Goodbye!") break elif cmd == "/clear": os.system('cls' if os.name == 'nt' else 'clear') continue elif cmd == "/temp" and len(parts) > 1: TEMPERATURE = float(parts[1]) print(f"Temperature set to {TEMPERATURE}") continue elif cmd == "/max" and len(parts) > 1: MAX_NEW_TOKENS = int(parts[1]) print(f"Max tokens set to {MAX_NEW_TOKENS}") continue else: print("Unknown command") continue if not prompt: print("Please enter a prompt!") continue # Generate print("\nGenerating...") result = generate(prompt) # Display print("\n" + "=" * 60) print("GENERATED TEXT:") print("=" * 60) print(result) print("=" * 60) # Statistics generated_part = result[len(prompt):] new_tokens = len(tokenizer.encode(generated_part)) print(f"Generated {new_tokens} new tokens") except KeyboardInterrupt: print("\n\nInterrupted. Type /exit to quit.") except Exception as e: print(f"\nError: {e}") if __name__ == "__main__": parser = argparse.ArgumentParser(description="VDrontMoE Text Generation") parser.add_argument("--prompt", type=str, help="Single prompt mode (no chat)") parser.add_argument("--temperature", type=float, default=0.4, help="Temperature (default: 0.4)") parser.add_argument("--max_tokens", type=int, default=100, help="Max new tokens") parser.add_argument("--top_p", type=float, default=0.9, help="Nucleus sampling threshold") parser.add_argument("--top_k", type=int, default=50, help="Top-K filtering") parser.add_argument("--model_dir", type=str, default="./VDrontMoE-2m-5e") args = parser.parse_args() # Update parameters MODEL_DIR = args.model_dir TEMPERATURE = args.temperature MAX_NEW_TOKENS = args.max_tokens TOP_P = args.top_p TOP_K = args.top_k if args.prompt: # Single mode print(f"\nPrompt: {args.prompt}") print(f"Temperature: {TEMPERATURE} | Max tokens: {MAX_NEW_TOKENS}") print("\n" + "=" * 60) result = generate(args.prompt) print(result) print("=" * 60) else: # Interactive mode chat()