[lint] format code with black (#22)

This commit is contained in:
Starrick Liu
2025-09-11 18:54:25 +08:00
committed by GitHub
parent 23dc63a399
commit 86f70a3b08
+12 -5
View File
@@ -30,7 +30,12 @@ class VQAWrapper(object):
return model return model
def generate(self, image: Image.Image, text: str, **kwargs) -> str: def generate(self, image: Image.Image, text: str, **kwargs) -> str:
messages = [{"role": "user", "content": [{"type": "image"}, {"type": "text", "text": text}]}] messages = [
{
"role": "user",
"content": [{"type": "image"}, {"type": "text", "text": text}],
}
]
text_prompt = self.processor.apply_chat_template( text_prompt = self.processor.apply_chat_template(
messages, tokenize=False, add_generation_prompt=True messages, tokenize=False, add_generation_prompt=True
) )
@@ -38,18 +43,19 @@ class VQAWrapper(object):
inputs = {k: v.to(self.device) for k, v in inputs.items()} inputs = {k: v.to(self.device) for k, v in inputs.items()}
generation_params = { generation_params = {
"max_new_tokens": 1024, # default value, can be overridden by kwargs "max_new_tokens": 1024, # default value, can be overridden by kwargs
"do_sample": False, "do_sample": False,
"eos_token_id": self.processor.tokenizer.eos_token_id, "eos_token_id": self.processor.tokenizer.eos_token_id,
"pad_token_id": self.processor.tokenizer.pad_token_id, "pad_token_id": self.processor.tokenizer.pad_token_id,
**kwargs **kwargs,
} }
with torch.no_grad(): with torch.no_grad():
generated_ids = self.model.generate(**inputs, **generation_params) generated_ids = self.model.generate(**inputs, **generation_params)
generated_ids = [ generated_ids = [
output_ids[len(input_ids):] for input_ids, output_ids in zip(inputs['input_ids'], generated_ids) output_ids[len(input_ids) :]
for input_ids, output_ids in zip(inputs["input_ids"], generated_ids)
] ]
response = self.processor.batch_decode( response = self.processor.batch_decode(
generated_ids, skip_special_tokens=True, clean_up_tokenization_spaces=False generated_ids, skip_special_tokens=True, clean_up_tokenization_spaces=False
@@ -67,9 +73,10 @@ if __name__ == "__main__":
# img = Image.open("/path/to/your/local/image.jpg").convert("RGB") # img = Image.open("/path/to/your/local/image.jpg").convert("RGB")
import requests import requests
img = Image.open(requests.get(test_image_url, stream=True).raw).convert("RGB") img = Image.open(requests.get(test_image_url, stream=True).raw).convert("RGB")
answer = wrapper.generate(img, test_question) answer = wrapper.generate(img, test_question)
print("model answer:", answer) print("model answer:", answer)
except Exception as e: except Exception as e:
print(f"model answer fail: {e}") print(f"model answer fail: {e}")