Views
No views yet


1import numpy as np
2import pandas as pd
3import torch
4import matplotlib.pyplot as plt
5from PIL import Image
6from datasets import load_dataset
7from torch.utils.data import Dataset
8from transformers import AutoImageProcessor, AutoTokenizer, VisionEncoderDecoderModel
91from datasets import load_dataset
2
3dataset = load_dataset("Seeker38/augmented_vi_face_wiki", split="train")1from transformers import AutoImageProcessor, AutoTokenizer, VisionEncoderDecoderModel
2model = VisionEncoderDecoderModel.from_pretrained("Seeker38/ViT_PhoBert_face_vi_wiki")
3phobert_tokenizer = AutoTokenizer.from_pretrained("vinai/phobert-base-v2", add_special_tokens=True)
4
5if phobert_tokenizer.pad_token is None:
6 phobert_tokenizer.add_special_tokens({'pad_token': '[PAD]'})1def generate_caption(model, dataset, tokenizer, device, num_images=20, max_length=50):
2 model.eval()
3
4 sampled_indices = random.sample(range(len(dataset)), num_images)
5 sampled_images = [dataset[idx]['image'] for idx in sampled_indices]
6 pixel_values_list = []
7
8 for image in sampled_images:
9 image = image.resize((224, 224))
10 image = np.array(image, dtype=np.uint8)
11 image = torch.tensor(np.moveaxis(image, -1, 0), dtype=torch.float32)
12 pixel_values_list.append(image)
13
14 pixel_values = torch.stack(pixel_values_list).to(device)
15
16 with torch.no_grad():
17 outputs = model.generate(pixel_values, num_beams=10, max_length=max_length, early_stopping=True, length_penalty=1.0)
18
19 decoded_preds = tokenizer.batch_decode(outputs, skip_special_tokens=True)
20
21 # Display the images and their captions in a single column
22 fig, axs = plt.subplots(num_images, 2, figsize=(15, 5 * num_images))
23
24 for i, (image, caption) in enumerate(zip(sampled_images, decoded_preds)):
25 axs[i, 0].imshow(image)
26 axs[i, 0].axis('off')
27 axs[i, 1].text(0, 0.5, caption, wrap=True, fontsize=12)
28 axs[i, 1].axis('off')
29
30 plt.tight_layout()
31
32 # Save the plot to a local file
33 output_file = "/kaggle/working/generated_captions.png"
34 plt.savefig(output_file)
35 plt.show()
36
37 print(f"Plot saved as {output_file}")generate_caption(model, dataset, phobert_tokenizer, device,5,70)