Views
No views yet
FlaxMarian.flax_clip_vision_marian folder in your project directory to load the model using the FlaxCLIPVisionMarianForConditionalGeneration class.1from torchvision.io import ImageReadMode, read_image
2from torchvision.transforms import CenterCrop, ConvertImageDtype, Normalize, Resize
3from torchvision.transforms.functional import InterpolationMode
4
5import torch
6import numpy as np
7from transformers import MarianTokenizer
8from flax_clip_vision_marian.modeling_clip_vision_marian import FlaxCLIPVisionMarianForConditionalGeneration
9
10clip_marian_model_name = 'flax-community/Image-captioning-Indonesia'
11model = FlaxCLIPVisionMarianForConditionalGeneration.from_pretrained(clip_marian_model_name)
12
13marian_model_name = 'Helsinki-NLP/opus-mt-en-id'
14tokenizer = MarianTokenizer.from_pretrained(marian_model_name)
15
16config = model.config
17image_size = config.clip_vision_config.image_size
18
19# Image transformation
20transforms = torch.nn.Sequential(
21 Resize([image_size], interpolation=InterpolationMode.BICUBIC),
22 CenterCrop(image_size),
23 ConvertImageDtype(torch.float),
24 Normalize((0.48145466, 0.4578275, 0.40821073), (0.26862954, 0.26130258, 0.27577711)),
25 )
26
27# Hyperparameters
28max_length = 8
29num_beams = 4
30gen_kwargs = {"max_length": max_length, "num_beams": num_beams}
31
32def generate_step(batch):
33 output_ids = model.generate(pixel_values, **gen_kwargs)
34 token_ids = np.array(output_ids.sequences)[0]
35 caption = tokenizer.decode(token_ids)
36 return caption
37
38image_file_path = image_file_path
39image = read_image(image_file_path, mode=ImageReadMode.RGB)
40image = transforms(image)
41pixel_values = torch.stack([image]).permute(0, 2, 3, 1).numpy()
42
43generated_ids = generate_step(pixel_values)
44
45print(generated_ids)