Image-Text-to-Text
Transformers
Safetensors
Polish
llama
text-generation
vision
multimodal
llava
siglip
bielik
llm
conversational
text-generation-inference
File size: 8,850 Bytes
fe45cbb
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
8350b14
fe45cbb
 
8350b14
fe45cbb
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
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<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 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()