stereoid's picture
Add files using upload-large-folder tool
1da285f verified
Raw
History Blame Contribute Delete
4.89 kB
import os
from PIL import Image
from transformers import AutoTokenizer
from vllm import LLM, SamplingParams
import ast
import time
import json
import logging
DATASET_PATH = '../../evaluation/gts/det/semantics/union3_test.json'
IMG_DIR = '../../dataset/data/coco_det/images/semantics/'
OUTPUT_PATH = './output.json'
model_path = "Aria-UI/Aria-UI-base"
MAX_ITER_PER_IMG = 60
PERFORM_ROUND = 5
llm = LLM(
model=model_path,
tokenizer_mode="slow",
dtype="bfloat16",
trust_remote_code=True,
gpu_memory_utilization=0.8,
enforce_eager=True,
tensor_parallel_size=2
)
tokenizer = AutoTokenizer.from_pretrained(
model_path, trust_remote_code=True, use_fast=False
)
def process_image(image_path, given_points = None):
given_points = given_points if given_points is not None else []
# instruction = 'Identify interactable objects.'
prompt_template = """Given a GUI image, what are the relative (0-1000) pixel point coordinates for the element corresponding to the following instruction or description: Identify all interactable objects in this VR scene.
You have already given the following pixel points before: {given_points}.
Do not identify duplicate interactable objects. If you have found all the interactable objects, respond with an empty list [].
"""
messages = [
{
"role": "user",
"content": [
{"type": "image"},
{
"type": "text",
"text": prompt_template.format(
given_points=given_points,
),
}
],
}
]
message = tokenizer.apply_chat_template(messages, add_generation_prompt=True)
outputs = llm.generate(
{
"prompt_token_ids": message,
"multi_modal_data": {
"image": [
Image.open(image_path),
],
"max_image_size": 980, # [Optional] The max image patch size, default `980`, maximum `980`, the image size for splitted blocks
"split_image": True, # [Optional] whether to split the images, default `True`
},
},
sampling_params=SamplingParams(max_tokens=50, top_k=1, stop=["<|im_end|>"]),
)
for o in outputs:
generated_tokens = o.outputs[0].token_ids
response = tokenizer.decode(generated_tokens, skip_special_tokens=True)
coords = ast.literal_eval(response.replace("<|im_end|>", "").replace("```", "").replace(" ", "").strip())
return coords
if __name__ == "__main__":
logging.basicConfig(format='%(asctime)s - %(levelname)s - %(name)s - %(message)s',
datefmt='%m/%d/%Y %H:%M:%S',level=logging.INFO )
start_time = time.time()
logging.info('Started at %s', time.strftime('%Y-%m-%d %H:%M:%S', time.localtime()))
with open(DATASET_PATH, 'r') as f:
dataset = json.load(f)
if os.path.exists(OUTPUT_PATH):
with open(OUTPUT_PATH, 'r') as f:
output = json.load(f)
else:
output = {}
# progress = tqdm(dataset['images'], total=len(dataset['images']), leave=True, position=0)
# for img in progress:
for img in dataset['images']:
img_id = img['id']
img_name = img['file_name']
img_path = os.path.join(IMG_DIR, img_name)
for r in range(PERFORM_ROUND):
if f"{img_id}_{r}" in output.keys() and not output[f"{img_id}_{r}"] == '':
logging.info(f"Skipping {img_name}_{r}")
continue
logging.info(f"Processing {img_name}")
# progress.set_description(f"Processing {img_name}")
objects = []
try:
for iter in range(MAX_ITER_PER_IMG):
coords = process_image(img_path, given_points=objects)
if coords in objects:
logging.info(f"Duplicate coordinates found. Ending iteration for {img_name}_{r}.")
break
if coords == []:
logging.info(f"No more objects found. Ending iteration for {img_name}_{r}.")
break
objects.append(coords)
except KeyboardInterrupt:
logging.info("Process interrupted by user. Exiting.")
break
except Exception as e:
logging.error(f"Error processing image {img_name}: {e}")
objects = ''
finally:
output[f"{img_id}_{r}"] = objects
with open(OUTPUT_PATH, 'w') as f:
json.dump(output, f, indent=4)
logging.info('Completed at %s', time.strftime('%Y-%m-%d %H:%M:%S', time.localtime()))
logging.info(f"Total time taken: {time.time() - start_time} seconds")