mihaimasala commited on
Commit
0cf879c
·
verified ·
1 Parent(s): 21f1159

Unify Getting Started snippet (consistency + decode slice fix)

Browse files
Files changed (1) hide show
  1. README.md +17 -10
README.md CHANGED
@@ -269,11 +269,13 @@ than Romanian.
269
  ```python
270
  import torch
271
  from PIL import Image
272
- from transformers import Qwen2VLForConditionalGeneration, AutoProcessor
273
 
274
  model = Qwen2VLForConditionalGeneration.from_pretrained(
275
- "OpenLLM-Ro/RoQwen2-VL-2B-Instruct", torch_dtype=torch.bfloat16, device_map="auto"
276
- )
 
 
277
  processor = AutoProcessor.from_pretrained("OpenLLM-Ro/RoQwen2-VL-2B-Instruct")
278
 
279
  image = Image.open("example.jpg").convert("RGB")
@@ -281,16 +283,21 @@ question = "Descrie imaginea în detaliu."
281
 
282
  messages = [
283
  {"role": "user", "content": [
284
- {"type": "image"},
285
  {"type": "text", "text": question},
286
  ]},
287
  ]
288
- text = processor.apply_chat_template(messages, tokenize=False, add_generation_prompt=True)
289
- inputs = processor(text=[text], images=[image], return_tensors="pt").to(model.device)
290
-
291
- outputs = model.generate(**inputs, max_new_tokens=256, do_sample=False)
292
- generated = outputs[:, inputs.input_ids.shape[1]:]
293
- print(processor.batch_decode(generated, skip_special_tokens=True)[0])
 
 
 
 
 
294
  ```
295
 
296
  ## Benchmarks
 
269
  ```python
270
  import torch
271
  from PIL import Image
272
+ from transformers import AutoProcessor, Qwen2VLForConditionalGeneration
273
 
274
  model = Qwen2VLForConditionalGeneration.from_pretrained(
275
+ "OpenLLM-Ro/RoQwen2-VL-2B-Instruct",
276
+ torch_dtype=torch.bfloat16,
277
+ device_map="auto",
278
+ ).eval()
279
  processor = AutoProcessor.from_pretrained("OpenLLM-Ro/RoQwen2-VL-2B-Instruct")
280
 
281
  image = Image.open("example.jpg").convert("RGB")
 
283
 
284
  messages = [
285
  {"role": "user", "content": [
286
+ {"type": "image", "image": image},
287
  {"type": "text", "text": question},
288
  ]},
289
  ]
290
+ inputs = processor.apply_chat_template(
291
+ messages,
292
+ add_generation_prompt=True,
293
+ tokenize=True,
294
+ return_dict=True,
295
+ return_tensors="pt",
296
+ ).to(model.device, dtype=torch.bfloat16)
297
+
298
+ with torch.inference_mode():
299
+ outputs = model.generate(**inputs, max_new_tokens=256, do_sample=False)
300
+ print(processor.decode(outputs[0][inputs["input_ids"].shape[-1]:], skip_special_tokens=True))
301
  ```
302
 
303
  ## Benchmarks