Views
No views yet
transformers):1import requests
2from PIL import Image
3import torch
4from transformers import AutoProcessor, Idefics3ForConditionalGeneration, TextIteratorStreamer, StoppingCriteria, StoppingCriteriaList
5
6base_model_id = "Andres77872/SmolVLM-500M-anime-caption-v0.2"
7
8processor = AutoProcessor.from_pretrained(base_model_id)
9model = Idefics3ForConditionalGeneration.from_pretrained(
10 base_model_id,
11 device_map="auto",
12 torch_dtype=torch.bfloat16
13)
14
15class StopOnTokens(StoppingCriteria):
16 def __init__(self, tokenizer, stop_sequence):
17 super().__init__()
18 self.tokenizer = tokenizer
19 self.stop_sequence = stop_sequence
20
21 def __call__(self, input_ids: torch.LongTensor, scores: torch.FloatTensor, **kwargs) -> bool:
22 new_text = self.tokenizer.decode(input_ids[0], skip_special_tokens=True)
23 max_keep = len(self.stop_sequence) + 10
24 if len(new_text) > max_keep:
25 new_text = new_text[-max_keep:]
26 return self.stop_sequence in new_text
27
28def prepare_inputs(image: Image.Image):
29 # IMPORTANT: The question prompt must remain fixed as "describe the image".
30 # This model is NOT designed for visual question answering.
31 # It is strictly an image captioning model, not intended to answer arbitrary visual questions.
32 question = "describe the image"
33 messages = [
34 {
35 "role": "user",
36 "content": [
37 {"type": "image"},
38 {"type": "text", "text": question}
39 ]
40 }
41 ]
42 max_image_size = processor.image_processor.max_image_size["longest_edge"]
43 size = processor.image_processor.size.copy()
44 if "longest_edge" in size and size["longest_edge"] > max_image_size:
45 size["longest_edge"] = max_image_size
46 prompt = processor.apply_chat_template(messages, add_generation_prompt=True)
47 inputs = processor(text=[prompt], images=[[image]], return_tensors='pt', padding=True, size=size)
48 inputs = {k: v.to(model.device) for k, v in inputs.items()}
49 return inputs
50
51# Example: caption a sample anime image
52image = Image.open(requests.get('https://img.arz.ai/5A7A-ckt', stream=True).raw).convert("RGB")
53inputs = prepare_inputs(image)
54stop_sequence = "</RATING>"
55streamer = TextIteratorStreamer(
56 processor.tokenizer,
57 skip_prompt=True,
58 skip_special_tokens=True,
59)
60custom_stopping_criteria = StoppingCriteriaList([
61 StopOnTokens(processor.tokenizer, stop_sequence)
62])
63
64with torch.no_grad():
65 generation_kwargs = dict(
66 **inputs,
67 streamer=streamer,
68 do_sample=False,
69 max_new_tokens=512,
70 pad_token_id=processor.tokenizer.pad_token_id,
71 stopping_criteria=custom_stopping_criteria,
72 )
73
74 import threading
75 generation_thread = threading.Thread(target=model.generate, kwargs=generation_kwargs)
76 generation_thread.start()
77
78 for new_text in streamer:
79 print(new_text, end="", flush=True)
80
81 generation_thread.join()