Image-Text-to-Text
Transformers
Safetensors
Polish
llama
text-generation
vision
multimodal
llava
siglip
bielik
llm
conversational
text-generation-inference
Bielik-1.5B-v3.0-VLM-Instruct / uruchom1_fullGUI.py
Wojtekb30's picture
Move model files to root directory
8350b14
Raw
History Blame Contribute Delete
14.7 kB
import os
import threading
import queue
import tkinter as tk
from tkinter import filedialog, messagebox, scrolledtext
from PIL import Image, ImageTk
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer, AutoModel, AutoImageProcessor
from safetensors.torch import load_file
# ============================================================
# KONFIGURACJA MODELI I URZĄDZENIA
# ============================================================
# Ścieżki / nazwy modeli
VISION_MODEL_PATH = "google/siglip-so400m-patch14-384"
#VISION_MODEL_PATH = "siglip-so400m-patch14-384"
MERGED_MODEL_PATH = "./"
PROJECTOR_FILE = "./mm_projector.safetensors"
DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
# ============================================================
# DEFINICJA PROJEKTORA MULTIMODALNEGO (TA SAMA CO WCZEŚNIEJ)
# ============================================================
class MultimodalProjector(torch.nn.Module):
def __init__(self, vision_dim, llm_dim):
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):
return self.net(x)
# ============================================================
# FUNKCJA ŁADUJĄCA MODELE (ZWRACA KOMPONENTY)
# ============================================================
def load_models(status_callback=None):
"""
Ładuje model językowy, tokenizer, vision tower, image processor i projektor.
status_callback - funkcja (str) -> None do aktualizowania statusu GUI.
Zwraca (model, vision_tower, vision_processor, projector, tokenizer)
"""
def stat(msg):
if status_callback:
status_callback(msg)
else:
print(msg)
stat(f"Ładowanie scalonego modelu z {MERGED_MODEL_PATH}...")
# 1. 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. Vision tower + processor
stat("Ładowanie Vision Tower...")
vision_tower = AutoModel.from_pretrained(
VISION_MODEL_PATH,
torch_dtype=torch.float16
).to(DEVICE)
vision_processor = AutoImageProcessor.from_pretrained(VISION_MODEL_PATH)
# 3. Projektor
stat("Ładowanie wag projektora...")
projector_weights = load_file(PROJECTOR_FILE)
vision_dim = vision_tower.config.vision_config.hidden_size
llm_dim = model.config.hidden_size
stat(f"Wymiary: Vision={vision_dim}, LLM={llm_dim}")
projector = MultimodalProjector(vision_dim=vision_dim, llm_dim=llm_dim).to(DEVICE).to(torch.float16)
projector.load_state_dict(projector_weights)
stat("Modele załadowane.")
return model, vision_tower, vision_processor, projector, tokenizer
# ============================================================
# LOGIKA GENEROWANIA ODPOWIEDZI (WYDZIELONE DO FUNKCJI)
# ============================================================
def generate_response(
model, vision_tower, vision_processor, projector, tokenizer,
prompt: str, image_pil: Image.Image | None, status_callback=None
):
"""
Generuje odpowiedź na zadany prompt. Jeśli image_pil jest None -> tryb tekstowy.
Zwraca wygenerowany tekst (str).
"""
def stat(msg):
if status_callback:
status_callback(msg)
else:
print(msg)
if model is None:
raise RuntimeError("Model językowy nie jest załadowany.")
if not prompt:
return ""
# Przygotuj obraz (jeśli jest)
pixel_values = None
has_image = False
if image_pil is not None:
try:
stat("Przetwarzanie obrazu...")
pixel_values = vision_processor(images=image_pil, return_tensors="pt").pixel_values
pixel_values = pixel_values.to(DEVICE, dtype=torch.float16)
has_image = True
except Exception as e:
raise RuntimeError(f"Błąd przetwarzania obrazu: {e}")
stat("Przygotowanie wejścia tekstowego...")
if has_image:
text_input = (
f"<|im_start|>user\n<image>\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 i embedowanie
input_ids = tokenizer.encode(text_input, return_tensors="pt").to(DEVICE)
inputs_embeds = model.model.embed_tokens(input_ids)
final_embeds = inputs_embeds
if has_image:
# ekstrakcja cech wizualnych i projektowanie
stat("Ekstrakcja cech wizualnych...")
with torch.no_grad():
vision_feats = vision_tower.vision_model(pixel_values).last_hidden_state # shape (1, seq_len_v, vision_dim)
img_embeds = projector(vision_feats)
# Tokenizacja clean (duplicated but consistent with two-stage in original)
input_ids_clean = tokenizer.encode(text_input, return_tensors="pt").to(DEVICE)
inputs_embeds_clean = model.model.embed_tokens(input_ids_clean)
final_embeds = torch.cat([img_embeds, inputs_embeds_clean], dim=1)
stat("Generowanie odpowiedzi (to może zająć chwilę)...")
with torch.no_grad():
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
)
generated_text = tokenizer.decode(output_ids[0], skip_special_tokens=True)
stat("Gotowe.")
return generated_text
# ============================================================
# GUI: APLIKACJA TKINTER
# ============================================================
class BielikGUI:
def __init__(self, root):
self.root = root
self.root.title("BIELIK LLaVA - GUI")
self.root.geometry("900x600")
# model components (załadowane przez load_models)
self.model = None
self.vision_tower = None
self.vision_processor = None
self.projector = None
self.tokenizer = None
# image state
self.image_path = None
self.image_pil = None
self.image_tk = None # referencja PhotoImage
# thread-safe queue to post status messages from worker threads to the UI
self.status_queue = queue.Queue()
self._build_ui()
# poll queue for status updates
self.root.after(100, self._poll_status_queue)
def _build_ui(self):
# Top frame: controls
ctrl_frame = tk.Frame(self.root)
ctrl_frame.pack(fill=tk.X, padx=8, pady=6)
self.load_models_btn = tk.Button(ctrl_frame, text="Load Models", command=self._on_load_models)
self.load_models_btn.pack(side=tk.LEFT, padx=4)
self.load_image_btn = tk.Button(ctrl_frame, text="Load Image", command=self._on_load_image)
self.load_image_btn.pack(side=tk.LEFT, padx=4)
self.remove_image_btn = tk.Button(ctrl_frame, text="Remove Image", command=self._on_remove_image, state=tk.DISABLED)
self.remove_image_btn.pack(side=tk.LEFT, padx=4)
self.generate_btn = tk.Button(ctrl_frame, text="Generate", command=self._on_generate, state=tk.DISABLED)
self.generate_btn.pack(side=tk.RIGHT, padx=4)
# Middle frame: image preview + prompt/response
middle_frame = tk.Frame(self.root)
middle_frame.pack(fill=tk.BOTH, expand=True, padx=8, pady=6)
# Left: image preview
preview_frame = tk.LabelFrame(middle_frame, text="Image Preview", width=300, height=300)
preview_frame.pack(side=tk.LEFT, fill=tk.BOTH, padx=6, pady=6)
preview_frame.pack_propagate(False)
self.image_label = tk.Label(preview_frame, text="No image loaded", anchor="center")
self.image_label.pack(fill=tk.BOTH, expand=True)
# Right: prompt & response
right_frame = tk.Frame(middle_frame)
right_frame.pack(side=tk.LEFT, fill=tk.BOTH, expand=True, padx=6, pady=6)
prompt_label = tk.Label(right_frame, text="Prompt:")
prompt_label.pack(anchor="w")
self.prompt_text = scrolledtext.ScrolledText(right_frame, height=6, wrap=tk.WORD)
self.prompt_text.pack(fill=tk.X, expand=False)
response_label = tk.Label(right_frame, text="Response:")
response_label.pack(anchor="w", pady=(8, 0))
self.response_text = scrolledtext.ScrolledText(right_frame, height=12, wrap=tk.WORD)
self.response_text.pack(fill=tk.BOTH, expand=True)
# Bottom: status bar
status_frame = tk.Frame(self.root)
status_frame.pack(fill=tk.X, padx=8, pady=4)
self.status_var = tk.StringVar(value="Ready.")
self.status_label = tk.Label(status_frame, textvariable=self.status_var, anchor="w")
self.status_label.pack(fill=tk.X)
# ------------------------
# Helpers for thread-safe status updates
# ------------------------
def _post_status(self, msg: str):
self.status_queue.put(msg)
def _poll_status_queue(self):
while not self.status_queue.empty():
try:
msg = self.status_queue.get_nowait()
except queue.Empty:
break
self.status_var.set(msg)
self.root.after(100, self._poll_status_queue)
# ------------------------
# GUI callbacks
# ------------------------
def _on_load_models(self):
# disable button while loading
self.load_models_btn.config(state=tk.DISABLED)
self._post_status("Starting model load...")
threading.Thread(target=self._load_models_worker, daemon=True).start()
def _load_models_worker(self):
try:
# call the previously defined loader with a status callback
def cb(msg): self._post_status(msg)
model, vision_tower, vision_processor, projector, tokenizer = load_models(status_callback=cb)
# set into self
self.model = model
self.vision_tower = vision_tower
self.vision_processor = vision_processor
self.projector = projector
self.tokenizer = tokenizer
self._post_status("Models loaded successfully.")
# enable the generate button
self.generate_btn.config(state=tk.NORMAL)
except Exception as e:
self._post_status(f"Error loading models: {e}")
messagebox.showerror("Model load error", str(e))
finally:
# ensure button re-enabled if load failed so user can retry
if self.model is None:
self.load_models_btn.config(state=tk.NORMAL)
def _on_load_image(self):
path = filedialog.askopenfilename(title="Wybierz obraz", filetypes=[("Image files", "*.png *.jpg *.jpeg *.bmp *.gif"), ("All files", "*.*")])
if not path:
return
try:
pil_img = Image.open(path).convert("RGB")
except Exception as e:
messagebox.showerror("Błąd", f"Nie można otworzyć pliku: {e}")
return
# create thumbnail for display
divider = pil_img.width / pil_img.height
max_size = (200, round(200 / divider))
img_preview = pil_img.copy()
self.image_tk = ImageTk.PhotoImage(img_preview.resize(max_size))
self.image_label.config(image=self.image_tk, text="")
self.image_path = path
self.image_pil = pil_img
self.remove_image_btn.config(state=tk.NORMAL)
self._post_status(f"Loaded image: {os.path.basename(path)}")
def _on_remove_image(self):
self.image_path = None
self.image_pil = None
self.image_tk = None
self.image_label.config(image="", text="No image loaded")
self.remove_image_btn.config(state=tk.DISABLED)
self._post_status("Image removed.")
def _on_generate(self):
if self.model is None:
messagebox.showwarning("Models not loaded", "Please load models first (press 'Load Models').")
return
prompt = self.prompt_text.get("1.0", tk.END).strip()
if not prompt:
messagebox.showinfo("No prompt", "Please enter a prompt.")
return
# disable generate button while running
self.generate_btn.config(state=tk.DISABLED)
self._post_status("Starting generation...")
threading.Thread(target=self._generate_worker, args=(prompt,), daemon=True).start()
def _generate_worker(self, prompt):
try:
result = generate_response(
model=self.model,
vision_tower=self.vision_tower,
vision_processor=self.vision_processor,
projector=self.projector,
tokenizer=self.tokenizer,
prompt=prompt,
image_pil=self.image_pil,
status_callback=lambda msg: self._post_status(msg)
)
# put result into response_text in main thread via queue
def put_result():
self.response_text.delete("1.0", tk.END)
self.response_text.insert(tk.END, result)
self._post_status("Generation finished.")
self.root.after(0, put_result)
except Exception as e:
self._post_status(f"Generation error: {e}")
messagebox.showerror("Generation error", str(e))
finally:
# re-enable generate button
def enable_btn():
self.generate_btn.config(state=tk.NORMAL)
self.root.after(0, enable_btn)
# ============================================================
# ENTRYPOINT
# ============================================================
def main():
root = tk.Tk()
app = BielikGUI(root)
root.mainloop()
if __name__ == "__main__":
main()