-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathmain_classification_with_vlm.py
More file actions
282 lines (214 loc) · 12.4 KB
/
Copy pathmain_classification_with_vlm.py
File metadata and controls
282 lines (214 loc) · 12.4 KB
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
272
273
274
275
276
277
278
279
280
281
from PIL import Image, ImageDraw, ImageFont
import torch
import base64
from openai import OpenAI
import pandas as pd
from tqdm import tqdm
import os
from sklearn.metrics import confusion_matrix
from PIL import Image, ImageDraw
import re
import cv2
import time
from app_config.settings import TEST_CROP_FILES_PATH, FONT_FILE, TEST_FULL_MODE_FILES_PATH, TRAIN_CROP_FILES, \
TRAIN_FULL_MODE_FILES_PATH
def encode_image(path):
with open(path, "rb") as f:
return base64.b64encode(f.read()).decode("utf-8")
def get_image_description_prompt(target_img, description_prompt):
# vllm
base64_img = encode_image(target_img)
messages = [
{
"role": "user",
"content": [
{"type": "text", "text": description_prompt},
{
"type": "image_url",
"image_url": {
"url": f"data:image/jpeg;base64,{base64_img}"},
},
],
}
]
return messages
def img_to_content(path):
return {
"type": "image_url",
"image_url": {
"url": f"data:image/jpeg;base64,{encode_image(path)}"
},
}
def get_classification_prompt(target_img_path, extract_descrtion):
#sa_22_path_1 = "/home/amitli/repo/dor6_vision/Dataset/few_shots/SA-22/11-21-02_1244400_1020.jpg"
sa_22_path_2 = "/home/amitli/repo/dor6_vision/Dataset/few_shots/SA-22/11-20-27_844400_795.jpg"
sa_22_path_3 = "/home/amitli/repo/dor6_vision/Dataset/few_shots/SA-22/11-17-44_444400_23.jpg"
#sa_22_txt_1 = "An overhead view shows a large wheeled military vehicle with a long rectangular chassis and multiple axles. The vehicle carries a rear‑mounted box‑shaped module that occupies most of the vehicle length. The top profile is flat and angular, with no visible gun barrel or turret. The overall silhouette is elongated and truck‑like, indicating a vehicle designed to carry a launcher or payload rather than direct‑fire weapons"
sa_22_txt_2 = "The image shows a tracked with long, rectangular body with a low profile."
sa_22_txt_3 = "The image shows a multi-wheeled truck. The front section is a cab with a flat windshield and a narrow profile, with side mirrors. It appears to have six wheels arranged in pairs along the chassis. The rear section has a raised, rectangular launch system holding multiple long, missile-like objects arranged in rows. There also appear to be exhaust pipes on the sides"
scud_path_1 = "/home/amitli/repo/dor6_vision/Dataset/few_shots/SCUD/11-21-10_1324400_673.jpg"
scud_path_2 = "/home/amitli/repo/dor6_vision/Dataset/few_shots/SCUD/11-17-54_524400_136.jpg"
scud_path_3 = "/home/amitli/repo/dor6_vision/Dataset/few_shots/SCUD/11-20-34_884400_1224.jpg"
scud_txt_1 = "The image shows a long, rectangular vehicle with a multi-wheeled chassis. It has a cab at the front and a long, enclosed cargo or equipment section behind it."
scud_txt_2 = "The image shows a long, rectangular vehicle with a cab at the front. A long cylindrical object (likely a missile) is mounted on top, extending significantly beyond the cab. The vehicle has multiple wheels—at least six are visible—arranged in pairs along its length"
scud_txt_3 = "The image shows a long, rectangular vehicle with a cylindrical object (likely a missile) mounted on top."
t_90_path_1 = "/home/amitli/repo/dor6_vision/Dataset/few_shots/T-90/11-20-40_924400_731.jpg"
#t_90_path_2 = "/home/amitli/repo/dor6_vision/Dataset/few_shots/T-90/11-18-04_484400_633.jpg"
t_90_path_3 = "/home/amitli/repo/dor6_vision/Dataset/few_shots/T-90/11-21-17_1284400_311.jpg"
t_90_txt_1 = "The image features a modern main battle tank with a distinct low-profile, rounded (hemispherical) turret. Centered in the turret is a long smoothbore main gun, often featuring a thermal sleeve or fume extractor. The hull is elongated and sits low to the ground, protected by thick frontal glacis armor, with heavy side skirts covering the upper portion of the tracks."
#t_90_txt_2 = "A main battle tank featuring a large, boxy turret with sharp, angled surfaces designed for kinetic energy deflection. It is armed with a large-caliber smoothbore gun. The exterior is notable for modular or composite armor plates bolted onto the turret faces and hull, creating a multi-layered, geometric appearance compared to cast-steel designs."
t_90_txt_3 = "A high-angle view of a main battle tank characterized by its continuous caterpillar tracks and a heavy armored hull. The most prominent feature is the centrally mounted, 360-degree rotating turret which houses a long-barrelled primary cannon. The silhouette is defined by the mechanical complexity of the drive sprockets and the low-profile chassis."
# VLLM:
messages = [
{
"role": "user",
"content": [
{"type": "text", "text": "You receive an image from a simulation and you must classify the military vehicle in the image, based only on the vehicle structure and mounted weapon system. Ignore background, terrain, and camera angle."},
{"type": "text", "text": "Identify the object in the frame. If the object is small or distant, consider its overall shape, color patterns."},
{"type": "text", "text": "Examples:"},
# Shot 1
#img_to_content(sa_22_path_1),
#{"type": "text", "text": sa_22_txt_1+ "\nAnswer: class_1"},
img_to_content(sa_22_path_2),
{"type": "text", "text": sa_22_txt_2 + "\nAnswer: class_1"},
img_to_content(sa_22_path_3),
{"type": "text", "text": sa_22_txt_3 + "\nAnswer: class_1"},
# Shot 2
img_to_content(scud_path_1),
{"type": "text", "text": scud_txt_1+ "\nAnswer: class_2"},
img_to_content(scud_path_2),
{"type": "text", "text": scud_txt_2 + "\nAnswer: class_2"},
img_to_content(scud_path_3),
{"type": "text", "text": scud_txt_3 + "\nAnswer: class_2"},
# Shot 3
img_to_content(t_90_path_1),
{"type": "text", "text": t_90_txt_1+ "\nAnswer: class_3"},
# img_to_content(t_90_path_2),
# {"type": "text", "text": t_90_txt_2+ "\nAnswer: class_3"},
img_to_content(t_90_path_3),
{"type": "text", "text": t_90_txt_3+ "\nAnswer: class_3"},
# Query image
img_to_content(target_img_path),
# {"type": "text",
# "text": "Analyze the provided image. First, describe the primary object's shape and visible features. Second, based on those features, classify the object into one of the following categories"},
{"type": "text",
# "text": "Based on the examples above, which class does this image belong to? Answer only: 'class_1', 'class_2', 'class_3' or 'Nothing'."}
"text": "Based on the examples above, which class does this image belong to? If the image does not fit any of the three, answer 'none'. Answer only: 'class_1', 'class_2', 'class_3', or 'none'."}
]
}
]
return messages
def plot_img_with_run_classification(image_path, classifcation):
img = Image.open(image_path).convert("RGB")
draw = ImageDraw.Draw(img)
if os.path.exists(FONT_FILE):
font = ImageFont.truetype(FONT_FILE, size=24)
else:
font = ImageFont.load_default()
draw.text(( 15, 15), classifcation, fill="red", font= font)
img.show()
def send_to_vllm(client, prompt_func, image_path, external_prompt=None):
try:
messages = prompt_func(image_path, external_prompt)
response = client.chat.completions.create(
#model="/model_path",
model = "google/gemma-4-31B-it",
messages=messages,
extra_body={
"mm_processor_kwargs": {
"max_soft_tokens": 1120
}
}
)
res_text = response.choices[0].message.content
except Exception as e:
print(f"Got an error: {e} FILE: {image_path}")
res_text = "Error"
return res_text
def run_train_classifcation(client, files_path, df):
d_convert = {"class_1": "SA-22",
"class_2": "SCUD",
"class_3": "T-90"}
l_jpg_file = []
l_gt = []
l_prediction = []
for i in tqdm(range(len(df))):
jpg_file = df.jpg_file.values[i]
gt = df['gt'].values[i]
full_file_path = f"{files_path}{jpg_file}"
classification = send_to_vllm(client, get_classification_prompt, full_file_path)
classification = classification.strip()
if classification.find('Answer:') != -1:
classification = classification[7:].strip()
if classification not in d_convert.keys():
#print(f"{jpg_file} = {classification}")
None
else:
classification = d_convert[classification]
l_jpg_file .append(jpg_file)
l_gt .append(gt)
l_prediction.append(classification)
df_res = pd.DataFrame({"jpg_file": l_jpg_file, "gt": l_gt, "prediction": l_prediction})
df_res.to_csv('/home/amitli/repo/dor6_vision/results/train_crop_vlm_classification.csv', index=False)
print_cm(df_res)
def print_cm(df):
classes = sorted(df['gt'].unique())
# compute confusion matrix (counts)
cm = confusion_matrix(df['gt'], df['prediction'], labels=classes)
# convert to DataFrame for readability
cm_df = pd.DataFrame(cm, index=classes, columns=classes)
# print("Confusion Matrix (Counts):")
# print(cm_df)
cm_percent = cm_df.div(cm_df.sum(axis=1), axis=0) * 100
print("\nConfusion Matrix (Percentages):")
print(cm_percent.round(2))
def eda_few_shots(client):
import glob
l_sa_22 = glob.glob('/home/amitli/repo/dor6_vision/Dataset/few_shots/SA-22/*.jpg')
l_scud = glob.glob('/home/amitli/repo/dor6_vision/Dataset/few_shots/SCUD/*.jpg')
l_t_90 = glob.glob('/home/amitli/repo/dor6_vision/Dataset/few_shots/T-90/*.jpg')
l_all = l_sa_22 + l_scud + l_t_90
l_gt = ['SA-22'] * 3 + ['SCUD'] * 3 + ['T-90'] * 3
for i in range(len(l_all)):
print("\n----------------------------------------------------------------------\n")
file = l_all[i]
gt = l_gt[i]
desc_prompt = f"You are getting a picture from a simulation that contains a {gt} military vehicle. Describe ONLY the visual features of the military vehicle in the picture (Only the visual features you see in this image)."
description = send_to_vllm(client, get_image_description_prompt, file, desc_prompt)
print(f"[{gt}] {os.path.basename(file)} = {description}")
def load_shiry_df():
df = pd.read_csv('/home/amitli/repo/dor6_vision/Dataset/shiry_testset_balanced.csv')
df = df.rename(columns={'filename': 'jpg_file', 'label_name': 'gt'})
return df
if __name__ == "__main__":
RUN_TRAIN = True
RUN_ON_TEST_SET = False
client = OpenAI(api_key="EMPTY", base_url="http://localhost:9000/v1")
# eda_few_shots(client)
# exit(0)
if RUN_TRAIN:
#FOLDER_PATH = TRAIN_CROP_FILES
FOLDER_PATH = TRAIN_FULL_MODE_FILES_PATH
#df = pd.read_csv('/home/amitli/repo/dor6_vision/Dataset/embeddings_train_crop.csv')
#df = df.sample(frac=0.10)
df = load_shiry_df()
run_train_classifcation(client, FOLDER_PATH, df)
exit(0)
if RUN_ON_TEST_SET:
df_test_crop = pd.read_csv('/home/amitli/repo/dor6_vision/Dataset/test_set_point.csv')
for i in range(len(df_test_crop)):
#i = 100
file = df_test_crop['jpg_file'].values[i]
full_crop_file = f"{TEST_CROP_FILES_PATH}{file}"
full_size_file = f"{TEST_FULL_MODE_FILES_PATH}{file}"
d_convert = {"Class_1": "SA-22",
"Class_2": "SCUD",
"Class_3": "T-90"}
# description = send_to_vllm(client, get_image_description_prompt, full_crop_file)
# print(description)
start_time = time.time()
classification = send_to_vllm(client, get_classification_prompt, full_crop_file)
end_time = time.time()
print(f"[{file}] time = {(end_time - start_time):.2f}s")
plot_img_with_run_classification(full_size_file, d_convert[classification])
#print(f"file = {file}")