import torch import os from PIL import Image from transformers import AutoModelForCausalLM, AutoTokenizer, AutoModel, AutoImageProcessor import tkinter as tk from tkinter import filedialog from safetensors.torch import load_file # ============================================================ # KONFIGURACJA ŚCIEŻEK MODELI ORAZ URZĄDZENIA OBLICZENIOWEGO # ============================================================ # Ścieżka do modelu wizyjnego VISION_MODEL_PATH = "google/siglip-so400m-patch14-384" #VISION_MODEL_PATH = "siglip-so400m-patch14-384" # Ścieżka do scalonego modelu językowego (LLM + adaptery) MERGED_MODEL_PATH = "./" # Ścieżka do pliku z wagami projektora multimodalnego PROJECTOR_FILE = "./mm_projector.safetensors" # Automatyczny wybór urządzenia: GPU (CUDA) jeśli dostępne, w przeciwnym razie CPU DEVICE = "cuda" if torch.cuda.is_available() else "cpu" def wybierz_plik(): """ Otwiera systemowe okno dialogowe wyboru pliku. Zwraca ścieżkę do wybranego pliku lub pusty string, jeśli użytkownik anulował wybór. """ root = tk.Tk() root.withdraw() sciezka = filedialog.askopenfilename() root.destroy() # Zamknięcie instancji Tkinter w celu zwolnienia zasobów return sciezka if sciezka else "" # ============================================================ # DEFINICJA PROJEKTORA MULTIMODALNEGO # ============================================================ class MultimodalProjector(torch.nn.Module): """ Projektor odpowiedzialny za mapowanie reprezentacji wizualnych (vision tower) do przestrzeni osadzeń (embeddingów) modelu językowego (LLM). Składa się z dwóch warstw liniowych z aktywacją GELU. """ def __init__(self, vision_dim, llm_dim): """ :param vision_dim: Rozmiar wektora cech z modelu vision. :param llm_dim: Rozmiar przestrzeni ukrytej modelu językowego. """ super().__init__() self.net = torch.nn.Sequential( torch.nn.Linear(vision_dim, llm_dim), torch.nn.GELU(), torch.nn.Linear(llm_dim, llm_dim) ) def forward(self, x): """ Przekształca wejściowe cechy wizualne do przestrzeni LLM. """ return self.net(x) def load_models(): """ Ładuje wszystkie wymagane komponenty systemu multimodalnego: - model językowy (LLM), - tokenizer, - model vision (vision tower), - procesor obrazu, - projektor multimodalny. Zwraca komplet załadowanych obiektów. """ print(f"Ładowanie Scalonego Bielika z {MERGED_MODEL_PATH}...") # 1. Ładowanie modelu językowego (LLM) model = AutoModelForCausalLM.from_pretrained( MERGED_MODEL_PATH, torch_dtype=torch.float16, trust_remote_code=True, device_map="auto" ) tokenizer = AutoTokenizer.from_pretrained(MERGED_MODEL_PATH) # 2. Ładowanie modelu vision oraz procesora obrazu print("Ładowanie Vision Tower...") vision_tower = AutoModel.from_pretrained( VISION_MODEL_PATH, torch_dtype=torch.float16 ).to(DEVICE) # Procesor odpowiada za preprocessing obrazu (resize, normalizacja itd.) vision_processor = AutoImageProcessor.from_pretrained(VISION_MODEL_PATH) # 3. Ładowanie wag projektora multimodalnego print("Ładowanie Projektora...") projector_weights = load_file(PROJECTOR_FILE) # Dynamiczne pobranie wymiarów ukrytych z konfiguracji modeli vision_dim = vision_tower.config.vision_config.hidden_size llm_dim = model.config.hidden_size print(f"Wykryte wymiary: Vision={vision_dim}, LLM={llm_dim}") # Inicjalizacja projektora z odpowiednimi wymiarami projector = MultimodalProjector( vision_dim=vision_dim, llm_dim=llm_dim ).to(DEVICE).to(torch.float16) # Załadowanie wag do projektora projector.load_state_dict(projector_weights) return model, vision_tower, vision_processor, projector, tokenizer def chat(): """ Główna pętla interakcyjna aplikacji. Umożliwia pracę w trybie: - tekstowym, - multimodalnym (tekst + obraz). """ # Załadowanie wszystkich komponentów systemu model, vision_tower, vision_processor, projector, tokenizer = load_models() print("\n" + "=" * 50) print("BIELIK LLaVA - GOTOWY DO ROZMOWY") print("Otworzy się okno wyboru pliku. Anulowanie wyboru = tryb tekstowy.") print("Wpisz 'exit' w konsoli aby wyjść.") print("=" * 50 + "\n") while True: # ============================================ # A. Wybór obrazu (opcjonalnie) # ============================================ print("\nWybierz plik obrazka w oknie...") img_path = wybierz_plik() pixel_values = None has_image = False if img_path: if os.path.exists(img_path): try: print(f"Wybrano: {img_path}") # Wczytanie i konwersja obrazu do RGB image = Image.open(img_path).convert('RGB') # Przetwarzanie obrazu do tensora zgodnego z vision tower pixel_values = vision_processor( images=image, return_tensors="pt" ).pixel_values.to(DEVICE, dtype=torch.float16) has_image = True except Exception as e: print(f"Błąd ładowania obrazka: {e}") continue else: print("Ścieżka nieprawidłowa.") else: print("Tryb tekstowy (bez obrazka).") # ============================================ # B. Pobranie promptu tekstowego # ============================================ prompt = input("Twój Prompt: ").strip() if prompt.lower() == 'exit': break if not prompt: continue # ============================================ # C. Generowanie odpowiedzi # ============================================ with torch.no_grad(): # 1. Ekstrakcja cech wizualnych (jeśli obraz został podany) img_embeds = None if has_image: vision_feats = vision_tower.vision_model(pixel_values).last_hidden_state img_embeds = projector(vision_feats) # 2. Przygotowanie wejścia tekstowego w odpowiednim formacie czatu if has_image: text_input = ( f"<|im_start|>user\n\n{prompt}<|im_end|>\n" f"<|im_start|>assistant\n" ) else: text_input = ( f"<|im_start|>user\n{prompt}<|im_end|>\n" f"<|im_start|>assistant\n" ) # Tokenizacja tekstu input_ids = tokenizer.encode( text_input, return_tensors="pt" ).to(DEVICE) # Konwersja tokenów do embeddingów inputs_embeds = model.model.embed_tokens(input_ids) # 3. Łączenie embeddingów obrazu i tekstu final_embeds = inputs_embeds if has_image: input_ids_clean = tokenizer.encode( text_input, return_tensors="pt" ).to(DEVICE) inputs_embeds_clean = model.model.embed_tokens(input_ids_clean) # Konkatenacja embeddingów obrazu oraz tekstu w osi sekwencji final_embeds = torch.cat( [img_embeds, inputs_embeds_clean], dim=1 ) # 4. Generowanie odpowiedzi przez model językowy print("Generowanie...", end="", flush=True) # Utworzenie attention_mask o tej samej długości co final_embeds batch_size = final_embeds.shape[0] seq_len = final_embeds.shape[1] attention_mask = torch.ones( (batch_size, seq_len), device=DEVICE ) output_ids = model.generate( inputs_embeds=final_embeds, attention_mask=attention_mask, max_new_tokens=256, temperature=0.3, do_sample=True, pad_token_id=tokenizer.eos_token_id, eos_token_id=tokenizer.eos_token_id ) # Dekodowanie wygenerowanej sekwencji tokenów generated_text = tokenizer.decode( output_ids[0], skip_special_tokens=True ) print("\r", end="") print(f"Bielik: {generated_text}") print("-" * 30) if __name__ == "__main__": chat()