Views
No views yet
1
2import argparse
3from pathlib import Path
4
5
6DEFAULT_SOURCE_MODEL = "google/gemma-4-E2B-it"
7
8
9def parse_args() -> argparse.Namespace:
10 parser = argparse.ArgumentParser(
11 description="Generate and optionally export a configurable tiny-random Gemma4 model.",
12 )
13 parser.add_argument("--output-dir", type=Path, required=True)
14 parser.add_argument("--source-model", default=DEFAULT_SOURCE_MODEL)
15 parser.add_argument("--hidden-size", type=int, default=32)
16 parser.add_argument("--head-dim", type=int, default=8)
17 parser.add_argument("--num-attention-heads", type=int, default=4)
18 parser.add_argument("--intermediate-size", type=int, default=64)
19 parser.add_argument("--text-layers", type=int, default=3)
20 parser.add_argument("--audio-layers", type=int, default=1)
21 parser.add_argument("--vision-layers", type=int, default=1)
22 parser.add_argument("--seed", type=int, default=0)
23 parser.add_argument("--export-dir", type=Path)
24 parser.add_argument("--smoke-test", action="store_true")
25 parser.add_argument("--device", default="CPU")
26 parser.add_argument("--attention-backend", choices=("PA", "SDPA"), default="SDPA")
27 return parser.parse_args()
28
29
30def validate_args(args: argparse.Namespace) -> None:
31 if args.hidden_size != args.head_dim * args.num_attention_heads:
32 raise ValueError(
33 "hidden-size must equal head-dim * num-attention-heads: "
34 f"{args.hidden_size} != {args.head_dim} * {args.num_attention_heads}"
35 )
36 if args.num_attention_heads < 2:
37 raise ValueError("num-attention-heads must be at least 2")
38 if args.text_layers != 3:
39 raise ValueError("text-layers must remain 3 for the model-card layer pattern")
40 if args.intermediate_size < args.hidden_size:
41 raise ValueError("intermediate-size must be at least hidden-size")
42 if args.smoke_test and args.export_dir is None:
43 raise ValueError("--smoke-test requires --export-dir")
44
45
46def generate_model(args: argparse.Namespace) -> None:
47 import torch
48 from transformers import AutoProcessor, Gemma4Config, Gemma4ForConditionalGeneration
49
50 torch.manual_seed(args.seed)
51 config = Gemma4Config.from_pretrained(args.source_model)
52
53 config.audio_config.hidden_size = args.hidden_size
54 config.audio_config.num_attention_heads = args.num_attention_heads
55 config.audio_config.num_hidden_layers = args.audio_layers
56 config.audio_config.output_proj_dims = args.hidden_size
57 config.audio_config.dtype = "float32"
58
59 config.text_config.global_head_dim = args.head_dim
60 config.text_config.head_dim = args.head_dim
61 config.text_config.hidden_size = args.hidden_size
62 config.text_config.hidden_size_per_layer_input = 1
63 config.text_config.intermediate_size = args.intermediate_size
64 config.text_config.num_attention_heads = args.num_attention_heads
65 config.text_config.num_key_value_heads = max(1, args.num_attention_heads // 2)
66 config.text_config.num_hidden_layers = args.text_layers
67 config.text_config.layer_types = ["sliding_attention", "full_attention", "full_attention"]
68 config.text_config.num_kv_shared_layers = 1
69 config.text_config.dtype = "float32"
70
71 config.vision_config.default_output_length = 70
72 config.vision_config.head_dim = args.head_dim
73 config.vision_config.hidden_size = args.hidden_size
74 config.vision_config.intermediate_size = args.intermediate_size
75 config.vision_config.num_attention_heads = args.num_attention_heads
76 config.vision_config.num_hidden_layers = args.vision_layers
77 config.vision_config.num_key_value_heads = args.num_attention_heads
78 config.vision_config.patch_size = 16
79 config.vision_config.dtype = "float32"
80
81 model = Gemma4ForConditionalGeneration(config)
82 model.eval()
83
84 args.output_dir.mkdir(parents=True, exist_ok=True)
85 model.save_pretrained(args.output_dir)
86 processor = AutoProcessor.from_pretrained(args.source_model, padding_side="left", truncation_side="left")
87 processor.save_pretrained(args.output_dir)
88
89 parameter_count = sum(parameter.numel() for parameter in model.parameters())
90 print(f"Saved {parameter_count:,}-parameter model to {args.output_dir}")
91
92 from transformers import AutoProcessor, Gemma4ForConditionalGeneration
93
94 messages = [
95 {
96 "role": "user", "content": [
97 {"type": "image",
98 "url": "https://raw.githubusercontent.com/google-gemma/cookbook/refs/heads/main/apps/sample-data/GoldenGate.png"},
99 {"type": "text", "text": "What is shown in this image?"}
100 ]
101 }
102 ]
103
104 processor = AutoProcessor.from_pretrained("google/gemma-4-E2B-it")
105 model = Gemma4ForConditionalGeneration.from_pretrained(
106 args.output_dir,
107 dtype="auto",
108 device_map="auto"
109 )
110
111 # Process input
112 inputs = processor.apply_chat_template(
113 messages,
114 tokenize=True,
115 return_dict=True,
116 return_tensors="pt",
117 add_generation_prompt=True,
118 ).to(model.device)
119 input_len = inputs["input_ids"].shape[-1]
120
121 # Generate output
122 outputs = model.generate(**inputs, max_new_tokens=512)
123 print("VLM infer OK")
124
125
126def export_model(model_dir: Path, export_dir: Path) -> None:
127 import openvino
128 import openvino_tokenizers
129 from optimum.intel.openvino import OVModelForVisualCausalLM
130 from transformers import AutoProcessor, Gemma4ForConditionalGeneration
131
132
133
134 processor = AutoProcessor.from_pretrained(model_dir, padding_side="left", truncation_side="left")
135 ov_model = OVModelForVisualCausalLM.from_pretrained(
136 model_dir,
137 compile=False,
138 device="CPU",
139 export=True,
140 load_in_8bit=False,
141 )
142
143 processor.image_processor.size ={
144 "height": 32,
145 "width": 32
146 }
147
148 export_dir.mkdir(parents=True, exist_ok=True)
149 ov_model.save_pretrained(export_dir)
150 processor.save_pretrained(export_dir)
151 ov_tokenizer, ov_detokenizer = openvino_tokenizers.convert_tokenizer(
152 processor.tokenizer,
153 with_detokenizer=True,
154 )
155 openvino.save_model(ov_tokenizer, export_dir / "openvino_tokenizer.xml")
156 openvino.save_model(ov_detokenizer, export_dir / "openvino_detokenizer.xml")
157 print(f"Exported OpenVINO model to {export_dir}")
158
159
160def smoke_test(export_dir: Path, device: str, attention_backend: str) -> None:
161 import numpy as np
162 import openvino
163 from openvino_genai import VLMPipeline
164 from transformers import Gemma4ForConditionalGeneration
165
166
167 pipeline = VLMPipeline(export_dir, device, ATTENTION_BACKEND=attention_backend)
168 text_result = pipeline.generate("Hello", max_new_tokens=1, do_sample=False)
169 print(f"Text smoke test passed: {text_result.texts!r}")
170
171 sampling_rate = 16_000
172 timestamps = np.arange(sampling_rate, dtype=np.float32) / sampling_rate
173 audio = 0.5 * np.sin(2 * np.pi * 440 * timestamps) + 0.25 * np.sin(2 * np.pi * 880 * timestamps)
174 audio_result = pipeline.generate(
175 "Describe this audio.<|audio|>",
176 audios=[openvino.Tensor(audio.astype(np.float32))],
177 max_new_tokens=1,
178 do_sample=False,
179 )
180 print(f"Audio smoke test passed: {audio_result.texts!r}")
181
182
183def main() -> None:
184 args = parse_args()
185 validate_args(args)
186 generate_model(args)
187 if args.export_dir is not None:
188 export_model(args.output_dir, args.export_dir)
189 if args.smoke_test:
190 smoke_test(args.export_dir, args.device, args.attention_backend)
191
192
193if __name__ == "__main__":
194 main()
195