Views
No views yet
1import os
2
3import accelerate
4import requests
5import torch
6import transformers
7from huggingface_hub import create_repo, upload_folder
8from PIL import Image
9from transformers import AutoProcessor, MllamaForConditionalGeneration
10from transformers.models.mllama import MllamaConfig
11
12model_id = 'meta-llama/Llama-3.2-90B-Vision-Instruct'
13repo_id = 'yujiepan/llama-3.2-vision-tiny-random'
14save_path = f'/tmp/{repo_id}'
15
16os.system(f'rm -rf {save_path}')
17
18config = transformers.AutoConfig.from_pretrained(
19 model_id,
20 trust_remote_code=True,
21)
22config.text_config.hidden_size = 8
23config.text_config.intermediate_size = 16
24config.text_config.num_attention_heads = 2
25config.text_config.num_key_value_heads = 1
26config.text_config.num_hidden_layers = 2
27config.text_config.cross_attention_layers = [1]
28
29config.vision_config.attention_heads = 2
30config.vision_config.hidden_size = 8
31config.vision_config.intermediate_size = 16
32config.vision_config.intermediate_layers_indices = [0]
33config.vision_config.num_global_layers = 2
34config.vision_config.num_hidden_layers = 2
35config.vision_config.vision_output_dim = 16
36
37
38transformers.set_seed(42)
39model = MllamaForConditionalGeneration(config)
40model.generation_config = transformers.GenerationConfig.from_pretrained(
41 model_id)
42model = model.to(torch.bfloat16)
43
44transformers.set_seed(42)
45with torch.no_grad():
46 for p in model.parameters():
47 torch.nn.init.normal_(p)
48
49model.save_pretrained(save_path)
50
51processor = AutoProcessor.from_pretrained(model_id)
52processor.save_pretrained(save_path)
53
54url = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/0052a70beed5bf71b92610a43a52df6d286cd5f3/diffusers/rabbit.jpg"
55image = Image.open(requests.get(url, stream=True).raw)
56
57messages = [
58 {"role": "user", "content": [
59 {"type": "image"},
60 {"type": "text", "text": "If I had to write a haiku for this one, it would be: "}
61 ]}
62]
63input_text = processor.apply_chat_template(
64 messages, add_generation_prompt=True)
65inputs = processor(image, input_text, return_tensors="pt").to(model.device)
66
67output = model.generate(**inputs, max_new_tokens=30)
68print(processor.decode(output[0]))
69
70os.system(f'ls -alh {save_path}')
71# os.system(f'rm -rf {save_path}/model.safetensors')
72# create_repo(repo_id, exist_ok=True)
73# upload_folder(repo_id=repo_id, folder_path=save_path)