comdoleger commited on
Commit
9648b93
·
verified ·
1 Parent(s): 97f0976

Upload extensions_built_in/dataset_tools/tools/fuyu_utils.py with huggingface_hub

Browse files
extensions_built_in/dataset_tools/tools/fuyu_utils.py ADDED
@@ -0,0 +1,66 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from transformers import CLIPImageProcessor, BitsAndBytesConfig, AutoTokenizer
2
+
3
+ from .caption import default_long_prompt, default_short_prompt, default_replacements, clean_caption
4
+ import torch
5
+ from PIL import Image
6
+
7
+
8
+ class FuyuImageProcessor:
9
+ def __init__(self, device='cuda'):
10
+ from transformers import FuyuProcessor, FuyuForCausalLM
11
+ self.device = device
12
+ self.model: FuyuForCausalLM = None
13
+ self.processor: FuyuProcessor = None
14
+ self.dtype = torch.bfloat16
15
+ self.tokenizer: AutoTokenizer
16
+ self.is_loaded = False
17
+
18
+ def load_model(self):
19
+ from transformers import FuyuProcessor, FuyuForCausalLM
20
+ model_path = "adept/fuyu-8b"
21
+ kwargs = {"device_map": self.device}
22
+ kwargs['load_in_4bit'] = True
23
+ kwargs['quantization_config'] = BitsAndBytesConfig(
24
+ load_in_4bit=True,
25
+ bnb_4bit_compute_dtype=self.dtype,
26
+ bnb_4bit_use_double_quant=True,
27
+ bnb_4bit_quant_type='nf4'
28
+ )
29
+ self.processor = FuyuProcessor.from_pretrained(model_path)
30
+ self.model = FuyuForCausalLM.from_pretrained(model_path, low_cpu_mem_usage=True, **kwargs)
31
+ self.is_loaded = True
32
+
33
+ self.tokenizer = AutoTokenizer.from_pretrained(model_path)
34
+ self.model = FuyuForCausalLM.from_pretrained(model_path, torch_dtype=self.dtype, **kwargs)
35
+ self.processor = FuyuProcessor(image_processor=FuyuImageProcessor(), tokenizer=self.tokenizer)
36
+
37
+ def generate_caption(
38
+ self, image: Image,
39
+ prompt: str = default_long_prompt,
40
+ replacements=default_replacements,
41
+ max_new_tokens=512
42
+ ):
43
+ # prepare inputs for the model
44
+ # text_prompt = f"{prompt}\n"
45
+
46
+ # image = image.convert('RGB')
47
+ model_inputs = self.processor(text=prompt, images=[image])
48
+ model_inputs = {k: v.to(dtype=self.dtype if torch.is_floating_point(v) else v.dtype, device=self.device) for k, v in
49
+ model_inputs.items()}
50
+
51
+ generation_output = self.model.generate(**model_inputs, max_new_tokens=max_new_tokens)
52
+ prompt_len = model_inputs["input_ids"].shape[-1]
53
+ output = self.tokenizer.decode(generation_output[0][prompt_len:], skip_special_tokens=True)
54
+ output = clean_caption(output, replacements=replacements)
55
+ return output
56
+
57
+ # inputs = self.processor(text=text_prompt, images=image, return_tensors="pt")
58
+ # for k, v in inputs.items():
59
+ # inputs[k] = v.to(self.device)
60
+
61
+ # # autoregressively generate text
62
+ # generation_output = self.model.generate(**inputs, max_new_tokens=max_new_tokens)
63
+ # generation_text = self.processor.batch_decode(generation_output[:, -max_new_tokens:], skip_special_tokens=True)
64
+ # output = generation_text[0]
65
+ #
66
+ # return clean_caption(output, replacements=replacements)