import os import json import threading import pygame import torch import torch.nn.functional as F from transformers import GPT2TokenizerFast, GPT2Config from safetensors.torch import load_file from model import VDrontModel # ------------------------------------------------------------ # Paths / constants # ------------------------------------------------------------ MODEL_DIR = "./VDrontV3-Mini" USER_TOKEN = "<|user|>" ASSISTANT_TOKEN = "<|assistant|>" WINDOW_WIDTH = 900 WINDOW_HEIGHT = 700 FPS = 60 # ------------------------------------------------------------ # Theme colors # ------------------------------------------------------------ def get_theme_colors(theme: str): if theme == "dark": return { "background": (25, 25, 30), "surface": (38, 38, 46), "surface_alt": (50, 50, 60), "text": (230, 230, 235), "text_secondary": (160, 160, 170), "accent": (100, 140, 255), "accent_hover": (130, 165, 255), "border": (70, 70, 85), "input_bg": (35, 35, 45), "user_bubble": (70, 100, 200), "user_text": (255, 255, 255), "ai_bubble": (52, 52, 62), "ai_text": (230, 230, 235), "button": (50, 50, 60), "button_hover": (70, 70, 85), "danger": (200, 80, 80), "success": (80, 180, 120), } else: # light return { "background": (240, 240, 245), "surface": (255, 255, 255), "surface_alt": (230, 230, 235), "text": (30, 30, 35), "text_secondary": (100, 100, 110), "accent": (60, 90, 200), "accent_hover": (90, 120, 230), "border": (200, 200, 210), "input_bg": (245, 245, 250), "user_bubble": (100, 140, 240), "user_text": (255, 255, 255), "ai_bubble": (225, 225, 230), "ai_text": (30, 30, 35), "button": (220, 220, 225), "button_hover": (200, 200, 210), "danger": (200, 80, 80), "success": (80, 180, 120), } # ------------------------------------------------------------ # Button class # ------------------------------------------------------------ class Button: def __init__(self, rect, text, callback): self.rect = pygame.Rect(rect) self.text = text self.callback = callback self.hovered = False def handle_event(self, event): if event.type == pygame.MOUSEMOTION: self.hovered = self.rect.collidepoint(event.pos) elif event.type == pygame.MOUSEBUTTONDOWN and event.button == 1: if self.rect.collidepoint(event.pos): self.callback() def draw(self, surface, colors, font): bg = colors["button_hover"] if self.hovered else colors["button"] pygame.draw.rect(surface, bg, self.rect, border_radius=6) pygame.draw.rect(surface, colors["border"], self.rect, width=1, border_radius=6) text_surf = font.render(self.text, True, colors["text"]) text_rect = text_surf.get_rect(center=self.rect.center) surface.blit(text_surf, text_rect) # ------------------------------------------------------------ # Main application # ------------------------------------------------------------ class VDrontLauncher: def __init__(self): pygame.init() self.screen = pygame.display.set_mode((WINDOW_WIDTH, WINDOW_HEIGHT)) pygame.display.set_caption("VDrontV3-Launcher") self.clock = pygame.time.Clock() self.running = True # State self.theme = "dark" self.colors = get_theme_colors(self.theme) self.settings_open = False self.qualitative = False self.input_text = "" self.input_active = True self.messages = [] self.scroll_offset = 0 self.max_scroll = 0 self.generating = False self.gen_thread = None self.gen_done = False self.gen_result = None # Generation params (normal mode by default) self.params = { "temperature": 0.45, "max_new_tokens": 256, "repetition_penalty": 1.1, "top_k": 50, "output_version": 0, } # Fonts self.font = self._get_font(18) self.font_small = self._get_font(14) self.font_big = self._get_font(22) # Load model self._show_loading("Loading model...") self.tokenizer, self.model, self.device = self._load_model() self._show_loading("Ready") # UI elements self.top_buttons = [] self.send_button = None self._create_buttons() # -------------------------------------------------------- # Fonts # -------------------------------------------------------- @staticmethod def _get_font(size, bold=False): candidates = ["Arial", "DejaVu Sans", "Segoe UI", "Verdana", "Helvetica"] for name in candidates: path = pygame.font.match_font(name, bold=bold) if path: return pygame.font.Font(path, size) return pygame.font.Font(None, size) def _show_loading(self, text): self.screen.fill(self.colors["background"]) surf = self.font_big.render(text, True, self.colors["text"]) rect = surf.get_rect(center=self.screen.get_rect().center) self.screen.blit(surf, rect) pygame.display.flip() # -------------------------------------------------------- # Model loading # -------------------------------------------------------- def _load_model(self): device = torch.device("cuda" if torch.cuda.is_available() else "cpu") tokenizer = GPT2TokenizerFast.from_pretrained(MODEL_DIR) vocab_size = len(tokenizer) special_tokens = [USER_TOKEN, ASSISTANT_TOKEN] tokenizer.add_special_tokens({"additional_special_tokens": special_tokens}) with open(os.path.join(MODEL_DIR, "architecture.json")) as f: arch = json.load(f) config = GPT2Config( vocab_size=vocab_size, n_embd=arch["n_embd"], n_head=arch["n_head"], n_layer=arch["n_layer"], n_positions=arch["n_positions"], layer_norm_epsilon=1e-5, ) model = VDrontModel( config=config, expert_start=arch["expert_start"], expert_end=arch["expert_end"], output_index=arch["output_index"], num_experts=arch["num_experts"], num_output_versions=arch["num_output_versions"], ) state = load_file(os.path.join(MODEL_DIR, "model.safetensors")) model.load_state_dict(state) model.to(device) model.eval() # Resize embeddings if tokenizer was extended if model.embed_tokens.num_embeddings < len(tokenizer): old_embed = model.embed_tokens new_embed = torch.nn.Embedding(len(tokenizer), old_embed.embedding_dim).to(device) new_embed.weight.data[:old_embed.num_embeddings] = old_embed.weight.data.to(device) model.embed_tokens = new_embed old_lm_head = model.lm_head new_lm_head = torch.nn.Linear(old_lm_head.in_features, len(tokenizer), bias=False).to(device) new_lm_head.weight.data[:old_lm_head.out_features] = old_lm_head.weight.data.to(device) model.lm_head = new_lm_head model.config.vocab_size = len(tokenizer) return tokenizer, model, device # -------------------------------------------------------- # UI creation # -------------------------------------------------------- def _create_buttons(self): self.theme_button = Button((20, 10, 120, 30), "", self._toggle_theme) self.qualitative_button = Button((150, 10, 140, 30), "", self._toggle_qualitative) self.settings_button = Button((300, 10, 100, 30), "Settings", self._open_settings) self.clear_button = Button((410, 10, 80, 30), "Clear", self._clear_chat) self.top_buttons = [ self.theme_button, self.qualitative_button, self.settings_button, self.clear_button, ] self.send_button = Button((WINDOW_WIDTH - 120, WINDOW_HEIGHT - 60, 100, 40), "Send", self._send_message) # -------------------------------------------------------- # Button callbacks # -------------------------------------------------------- def _toggle_theme(self): self.theme = "light" if self.theme == "dark" else "dark" self.colors = get_theme_colors(self.theme) def _toggle_qualitative(self): self.qualitative = not self.qualitative if self.qualitative: self.params = { "temperature": 0.3, "max_new_tokens": 512, "repetition_penalty": 1.4, "top_k": 50, "output_version": 1, } else: self.params = { "temperature": 0.45, "max_new_tokens": 256, "repetition_penalty": 1.1, "top_k": 50, "output_version": 0, } def _open_settings(self): self.settings_open = True def _clear_chat(self): self.messages.clear() self.scroll_offset = 0 # -------------------------------------------------------- # Generation (run in separate thread) # -------------------------------------------------------- def _generate_thread(self, prompt): try: self.model.set_output_version(self.params["output_version"]) input_ids = self.tokenizer.encode(prompt, return_tensors="pt").to(self.device) generated_tokens = [] eos_id = self.tokenizer.eos_token_id with torch.no_grad(): for _ in range(self.params["max_new_tokens"]): pos = torch.arange(0, input_ids.size(1), device=self.device).unsqueeze(0) x = self.model.embed_tokens(input_ids) + self.model.embed_positions(pos) router_logits = self.model.router(x.mean(dim=1)) expert_idx = router_logits.argmax(dim=-1).item() self.model.set_expert_version(expert_idx) idx_cond = input_ids[:, -self.model.config.n_positions:] logits, _ = self.model(idx_cond) logits = logits[:, -1, :] / self.params["temperature"] for token_id in set(input_ids[0].tolist()): logits[0, token_id] /= self.params["repetition_penalty"] if self.params["top_k"] is not None and self.params["top_k"] > 0: v, _ = torch.topk(logits, min(self.params["top_k"], logits.size(-1))) logits[logits < v[:, [-1]]] = -float("Inf") probs = F.softmax(logits, dim=-1) idx_next = torch.multinomial(probs, num_samples=1) next_token = idx_next.item() if next_token == eos_id: break generated_tokens.append(next_token) input_ids = torch.cat((input_ids, idx_next), dim=1) full_text = self.tokenizer.decode(generated_tokens, skip_special_tokens=True).strip() self.gen_result = full_text except Exception as e: self.gen_result = f"[Error] {e}" finally: self.gen_done = True def _send_message(self): text = self.input_text.strip() if not text or self.generating: return self.messages.append({"role": "user", "text": text}) self.input_text = "" self.scroll_offset = 0 prompt = f"{USER_TOKEN}{text}{ASSISTANT_TOKEN}" self.generating = True self.gen_done = False self.gen_result = None self.gen_thread = threading.Thread(target=self._generate_thread, args=(prompt,), daemon=True) self.gen_thread.start() # -------------------------------------------------------- # Event handling # -------------------------------------------------------- def _handle_events(self): for event in pygame.event.get(): if event.type == pygame.QUIT: self.running = False if self.settings_open: self._handle_settings_event(event) else: self._handle_main_event(event) def _handle_main_event(self, event): # Buttons for btn in self.top_buttons: btn.handle_event(event) self.send_button.handle_event(event) # Mouse wheel scroll if event.type == pygame.MOUSEWHEEL: self.scroll_offset = max(0, min(self.max_scroll, self.scroll_offset - event.y * 30)) # Mouse click for input activation if event.type == pygame.MOUSEBUTTONDOWN and event.button == 1: input_rect = pygame.Rect(20, WINDOW_HEIGHT - 60, WINDOW_WIDTH - 140, 40) self.input_active = input_rect.collidepoint(event.pos) # Keyboard input if event.type == pygame.KEYDOWN: if event.key == pygame.K_RETURN: self._send_message() elif event.key == pygame.K_BACKSPACE: self.input_text = self.input_text[:-1] elif event.unicode and event.unicode.isprintable(): self.input_text += event.unicode def _handle_settings_event(self, event): if event.type == pygame.MOUSEBUTTONDOWN and event.button == 1: # Close button / click outside panel_rect = pygame.Rect( (WINDOW_WIDTH - 500) // 2, (WINDOW_HEIGHT - 380) // 2, 500, 380, ) close_rect = pygame.Rect(panel_rect.right - 35, panel_rect.y + 10, 25, 25) if close_rect.collidepoint(event.pos) or not panel_rect.collidepoint(event.pos): self.settings_open = False return # Check row controls for key, action, rect in self.settings_controls: if rect.collidepoint(event.pos): self._adjust_param(key, action) break # -------------------------------------------------------- # Settings # -------------------------------------------------------- def _adjust_param(self, key, action): row = next((r for r in self.settings_rows if r["key"] == key), None) if not row: return if key == "output_version": self.params[key] = 0 if self.params[key] == 1 else 1 else: step = row["step"] value = self.params[key] new_value = value + step if action == "plus" else value - step new_value = max(row["min"], min(row["max"], new_value)) if isinstance(step, int): new_value = int(round(new_value)) else: new_value = round(new_value, 2) self.params[key] = new_value # Manual adjustment means qualitative preset is no longer active self.qualitative = False # -------------------------------------------------------- # Update # -------------------------------------------------------- def _update(self): # Update button labels self.theme_button.text = f"Theme: {'Dark' if self.theme == 'dark' else 'Light'}" self.qualitative_button.text = f"Qualitative: {'ON' if self.qualitative else 'OFF'}" # Check generation completion if self.generating and self.gen_done: result = self.gen_result if self.gen_result is not None else "[No response]" self.messages.append({"role": "ai", "text": result}) self.generating = False self.gen_done = False self.gen_result = None self.gen_thread = None self.scroll_offset = 0 # -------------------------------------------------------- # Drawing # -------------------------------------------------------- def _draw(self): self.screen.fill(self.colors["background"]) self._draw_top_bar() self._draw_chat() self._draw_input() if self.generating: self._draw_typing_indicator() if self.settings_open: self._draw_settings() pygame.display.flip() def _draw_top_bar(self): for btn in self.top_buttons: btn.draw(self.screen, self.colors, self.font_small) def _draw_input(self): input_rect = pygame.Rect(20, WINDOW_HEIGHT - 60, WINDOW_WIDTH - 140, 40) pygame.draw.rect(self.screen, self.colors["input_bg"], input_rect, border_radius=6) pygame.draw.rect(self.screen, self.colors["border"], input_rect, width=1, border_radius=6) # Render input text (clipped) text_surf = self.font.render(self.input_text, True, self.colors["text"]) clip_rect = input_rect.inflate(-10, -10) self.screen.set_clip(clip_rect) self.screen.blit(text_surf, (input_rect.x + 10, input_rect.y + 8)) self.screen.set_clip(None) # Blinking cursor if self.input_active and pygame.time.get_ticks() % 1000 < 500: cursor_x = input_rect.x + 10 + text_surf.get_width() + 2 if cursor_x < input_rect.right - 10: pygame.draw.line( self.screen, self.colors["text"], (cursor_x, input_rect.y + 8), (cursor_x, input_rect.y + 32), 2, ) self.send_button.draw(self.screen, self.colors, self.font) def _draw_typing_indicator(self): text = "AI is typing..." surf = self.font_small.render(text, True, self.colors["text_secondary"]) rect = surf.get_rect(topleft=(20, WINDOW_HEIGHT - 75)) self.screen.blit(surf, rect) def _draw_chat(self): chat_rect = pygame.Rect(20, 50, WINDOW_WIDTH - 40, WINDOW_HEIGHT - 130) pygame.draw.rect(self.screen, self.colors["surface"], chat_rect, border_radius=8) # Calculate total content height for scrollbar total_height = 0 wrapped_cache = [] for msg in self.messages: bubble_width = chat_rect.width - 40 wrapped = self._wrap_text(msg["text"], self.font, bubble_width - 20) line_height = self.font.get_linesize() bubble_height = line_height * len(wrapped) + 20 total_height += bubble_height + 10 # spacing wrapped_cache.append((msg, wrapped, bubble_height)) self.max_scroll = max(0, total_height - chat_rect.height) self.scroll_offset = max(0, min(self.scroll_offset, self.max_scroll)) self.screen.set_clip(chat_rect) y = chat_rect.bottom - 10 + self.scroll_offset for msg, wrapped, bubble_height in reversed(wrapped_cache): bubble_rect = pygame.Rect(chat_rect.x + 10, y - bubble_height, chat_rect.width - 40, bubble_height) if bubble_rect.bottom < chat_rect.top: break if bubble_rect.top <= chat_rect.bottom: if msg["role"] == "user": bubble_rect.right = chat_rect.right - 10 bg = self.colors["user_bubble"] fg = self.colors["user_text"] else: bubble_rect.left = chat_rect.x + 10 bg = self.colors["ai_bubble"] fg = self.colors["ai_text"] pygame.draw.rect(self.screen, bg, bubble_rect, border_radius=10) line_height = self.font.get_linesize() text_y = bubble_rect.y + 10 for line in wrapped: line_surf = self.font.render(line, True, fg) if msg["role"] == "user": self.screen.blit(line_surf, (bubble_rect.right - 15 - line_surf.get_width(), text_y)) else: self.screen.blit(line_surf, (bubble_rect.x + 15, text_y)) text_y += line_height y = bubble_rect.y - 10 self.screen.set_clip(None) # Scrollbar if total_height > chat_rect.height: scrollbar_height = max(30, int(chat_rect.height * (chat_rect.height / total_height))) scrollbar_y = chat_rect.y + int((chat_rect.height - scrollbar_height) * (self.scroll_offset / self.max_scroll)) if self.max_scroll > 0 else chat_rect.y scrollbar_rect = pygame.Rect(chat_rect.right - 6, scrollbar_y, 4, scrollbar_height) pygame.draw.rect(self.screen, self.colors["border"], scrollbar_rect, border_radius=2) def _draw_settings(self): panel_width = 500 panel_height = 380 panel_x = (WINDOW_WIDTH - panel_width) // 2 panel_y = (WINDOW_HEIGHT - panel_height) // 2 panel_rect = pygame.Rect(panel_x, panel_y, panel_width, panel_height) # Overlay overlay = pygame.Surface((WINDOW_WIDTH, WINDOW_HEIGHT), pygame.SRCALPHA) overlay.fill((0, 0, 0, 128)) self.screen.blit(overlay, (0, 0)) pygame.draw.rect(self.screen, self.colors["surface"], panel_rect, border_radius=12) pygame.draw.rect(self.screen, self.colors["border"], panel_rect, width=2, border_radius=12) # Title title_surf = self.font_big.render("Settings", True, self.colors["text"]) self.screen.blit(title_surf, (panel_x + 20, panel_y + 15)) # Close button close_rect = pygame.Rect(panel_rect.right - 35, panel_y + 10, 25, 25) pygame.draw.rect(self.screen, self.colors["button"], close_rect, border_radius=6) pygame.draw.rect(self.screen, self.colors["border"], close_rect, width=1, border_radius=6) close_text = self.font_small.render("X", True, self.colors["text"]) self.screen.blit(close_text, close_text.get_rect(center=close_rect.center)) # Settings rows self.settings_rows = [ {"key": "temperature", "label": "Temperature", "min": 0.1, "max": 2.0, "step": 0.05}, {"key": "max_new_tokens", "label": "Max Tokens", "min": 32, "max": 1024, "step": 32}, {"key": "repetition_penalty", "label": "Repetition Penalty", "min": 0.8, "max": 2.0, "step": 0.1}, {"key": "top_k", "label": "Top K", "min": 0, "max": 100, "step": 5}, {"key": "output_version", "label": "Output Version", "min": 0, "max": 1, "step": 1}, ] self.settings_controls = [] for i, row in enumerate(self.settings_rows): y = panel_y + 70 + i * 55 # Label label_surf = self.font.render(row["label"], True, self.colors["text"]) self.screen.blit(label_surf, (panel_x + 25, y)) # Minus button minus_rect = pygame.Rect(panel_x + 310, y, 30, 30) pygame.draw.rect(self.screen, self.colors["button"], minus_rect, border_radius=6) pygame.draw.rect(self.screen, self.colors["border"], minus_rect, width=1, border_radius=6) minus_text = self.font.render("-", True, self.colors["text"]) self.screen.blit(minus_text, minus_text.get_rect(center=minus_rect.center)) self.settings_controls.append((row["key"], "minus", minus_rect)) # Value value_surf = self.font.render(str(self.params[row["key"]]), True, self.colors["text"]) value_rect = value_surf.get_rect(center=(panel_x + 370, y + 15)) self.screen.blit(value_surf, value_rect) # Plus button plus_rect = pygame.Rect(panel_x + 410, y, 30, 30) pygame.draw.rect(self.screen, self.colors["button"], plus_rect, border_radius=6) pygame.draw.rect(self.screen, self.colors["border"], plus_rect, width=1, border_radius=6) plus_text = self.font.render("+", True, self.colors["text"]) self.screen.blit(plus_text, plus_text.get_rect(center=plus_rect.center)) self.settings_controls.append((row["key"], "plus", plus_rect)) # -------------------------------------------------------- # Text wrapping # -------------------------------------------------------- def _wrap_text(self, text, font, max_width): words = text.split(" ") lines = [] current = "" for word in words: test = word if not current else current + " " + word if font.size(test)[0] <= max_width: current = test else: if current: lines.append(current) current = word else: # Very long word, split by characters while font.size(word)[0] > max_width: split_idx = len(word) for i in range(1, len(word)): if font.size(word[:i])[0] > max_width: split_idx = i - 1 break if split_idx == len(word): break lines.append(word[:split_idx]) word = word[split_idx:] current = word if current: lines.append(current) return lines # -------------------------------------------------------- # Main loop # -------------------------------------------------------- def run(self): while self.running: self.clock.tick(FPS) self._handle_events() self._update() self._draw() pygame.quit() if __name__ == "__main__": app = VDrontLauncher() app.run()