Views
No views yet
Helsinki-NLP-opus-mt-en-de/
├── config.json
├── decoder_model.onnx
├── decoder_model_merged.onnx
├── decoder_with_past_model.onnx
├── encoder_model.onnx
├── generation_config.json
├── source.spm
├── target.spm
├── special_tokens_map.json
├── tokenizer_config.json
└── vocab.jsonvocab.json, source.spm, target.spm, etc.pip install huggingface_hub onnxruntime transformers sentencepiece1"""
2ONNX Runtime MarianMT Translation Demo
3======================================
4
5This demo performs text translation using:
6- HuggingFace MarianTokenizer
7- ONNX Runtime Encoder
8- ONNX Runtime Merged Decoder with KV Cache
9
10The decoder uses past key/value attention caching to avoid
11recomputing previous tokens, making autoregressive generation
12significantly faster.
13
14Model:
15 Helsinki-NLP/opus-mt-en-de
16
17Task:
18 English → German translation
19"""
20
21
22import json
23import numpy as np
24import onnxruntime as ort
25
26from huggingface_hub import snapshot_download
27from transformers import MarianTokenizer
28
29
30# ============================================================
31# 1. Download Model Files
32# ============================================================
33
34print("\n[1/6] Downloading ONNX model...")
35
36model_root = snapshot_download(
37 repo_id="VaishalBusiness/opus",
38 allow_patterns="Helsinki-NLP-opus-mt-en-de/*",
39)
40
41model_dir = f"{model_root}/Helsinki-NLP-opus-mt-en-de"
42
43print(f"Model loaded from:\n{model_dir}")
44
45
46# ============================================================
47# 2. Load Tokenizer and Model Configuration
48# ============================================================
49
50print("\n[2/6] Loading tokenizer and configuration...")
51
52tokenizer = MarianTokenizer.from_pretrained(model_dir)
53
54
55with open(f"{model_dir}/config.json") as file:
56 config = json.load(file)
57
58
59with open(f"{model_dir}/generation_config.json") as file:
60 generation_config = json.load(file)
61
62
63# Generation parameters
64eos_token_id = generation_config.get(
65 "eos_token_id",
66 config.get("eos_token_id")
67)
68
69pad_token_id = generation_config.get(
70 "pad_token_id",
71 config.get("pad_token_id")
72)
73
74decoder_start_token_id = generation_config.get(
75 "decoder_start_token_id",
76 config.get(
77 "decoder_start_token_id",
78 pad_token_id
79 )
80)
81
82
83# Transformer architecture parameters
84num_layers = config["decoder_layers"]
85num_heads = config["decoder_attention_heads"]
86
87hidden_size = config["d_model"]
88head_dim = hidden_size // num_heads
89
90
91print(
92 f"""
93Transformer Configuration:
94--------------------------
95Layers : {num_layers}
96Attention heads : {num_heads}
97Head dimension : {head_dim}
98Hidden size : {hidden_size}
99"""
100)
101
102
103# ============================================================
104# 3. Tokenize Input Text
105# ============================================================
106
107print("\n[3/6] Encoding input text...")
108
109
110source_text = "Hello, how are you?"
111
112
113encoded = tokenizer(
114 source_text,
115 return_tensors="np"
116)
117
118
119input_ids = encoded["input_ids"].astype(np.int64)
120attention_mask = encoded["attention_mask"].astype(np.int64)
121
122
123print("Input tokens:")
124print(input_ids)
125
126
127# ============================================================
128# 4. Run Encoder
129# ============================================================
130
131print("\n[4/6] Running encoder...")
132
133
134encoder = ort.InferenceSession(
135 f"{model_dir}/encoder_model.onnx"
136)
137
138
139encoder_outputs = encoder.run(
140 None,
141 {
142 "input_ids": input_ids,
143 "attention_mask": attention_mask
144 }
145)
146
147
148encoder_hidden_states = encoder_outputs[0]
149
150
151print(
152 "Encoder output shape:",
153 encoder_hidden_states.shape
154)
155
156
157
158# ============================================================
159# 5. Initialize Decoder With KV Cache
160# ============================================================
161
162print("\n[5/6] Initializing decoder KV cache...")
163
164
165decoder = ort.InferenceSession(
166 f"{model_dir}/decoder_model_merged.onnx"
167)
168
169
170decoder_inputs = {
171 item.name
172 for item in decoder.get_inputs()
173}
174
175
176decoder_outputs = [
177 item.name
178 for item in decoder.get_outputs()
179]
180
181
182batch_size = encoder_hidden_states.shape[0]
183
184
185# Empty cache for first decoding step.
186# Shape:
187# [batch, heads, sequence_length, head_dimension]
188#
189# sequence_length = 0 because no tokens have been generated yet.
190
191empty_cache = np.zeros(
192 (
193 batch_size,
194 num_heads,
195 0,
196 head_dim
197 ),
198 dtype=np.float32
199)
200
201
202past_key_values = {}
203
204
205for layer in range(num_layers):
206
207 past_key_values[
208 f"past_key_values.{layer}.decoder.key"
209 ] = empty_cache
210
211 past_key_values[
212 f"past_key_values.{layer}.decoder.value"
213 ] = empty_cache
214
215
216 # Encoder KV cache is generated during the first decoder call
217 past_key_values[
218 f"past_key_values.{layer}.encoder.key"
219 ] = empty_cache
220
221 past_key_values[
222 f"past_key_values.{layer}.encoder.value"
223 ] = empty_cache
224
225
226
227# ============================================================
228# 6. Autoregressive Generation Loop
229# ============================================================
230
231print("\n[6/6] Generating translation...\n")
232
233
234generated_tokens = [
235 decoder_start_token_id
236]
237
238
239decoder_input_ids = np.array(
240 [[decoder_start_token_id]],
241 dtype=np.int64
242)
243
244
245# False = first decoder pass
246# True = reuse KV cache
247use_cache_branch = np.array(
248 [False],
249 dtype=bool
250)
251
252
253
254MAX_LENGTH = 128
255
256
257for step in range(MAX_LENGTH):
258
259
260 # Prepare decoder inputs
261 feed = {
262
263 "input_ids":
264 decoder_input_ids,
265
266 "encoder_hidden_states":
267 encoder_hidden_states,
268
269 "encoder_attention_mask":
270 attention_mask,
271
272 "use_cache_branch":
273 use_cache_branch,
274 }
275
276
277 # Add cached attention states
278 feed.update(
279 {
280 key: value
281 for key, value in past_key_values.items()
282 if key in decoder_inputs
283 }
284 )
285
286
287 # Run decoder
288 output = decoder.run(
289 decoder_outputs,
290 feed
291 )
292
293
294 output = dict(
295 zip(
296 decoder_outputs,
297 output
298 )
299 )
300
301
302 # Select highest probability token
303 logits = output["logits"]
304
305 next_token = int(
306 np.argmax(
307 logits[0, -1]
308 )
309 )
310
311
312 generated_tokens.append(
313 next_token
314 )
315
316
317 # Stop at EOS token
318 if next_token == eos_token_id:
319 break
320
321
322
323 # Update KV cache
324 updated_cache = dict(
325 past_key_values
326 )
327
328
329 for layer in range(num_layers):
330
331 updated_cache[
332 f"past_key_values.{layer}.decoder.key"
333 ] = output[
334 f"present.{layer}.decoder.key"
335 ]
336
337 updated_cache[
338 f"past_key_values.{layer}.decoder.value"
339 ] = output[
340 f"present.{layer}.decoder.value"
341 ]
342
343
344 # Encoder cache only needs to be stored once
345 if step == 0:
346
347 updated_cache[
348 f"past_key_values.{layer}.encoder.key"
349 ] = output[
350 f"present.{layer}.encoder.key"
351 ]
352
353 updated_cache[
354 f"past_key_values.{layer}.encoder.value"
355 ] = output[
356 f"present.{layer}.encoder.value"
357 ]
358
359
360 past_key_values = updated_cache
361
362
363 # Next step only feeds the newly generated token
364 decoder_input_ids = np.array(
365 [[next_token]],
366 dtype=np.int64
367 )
368
369
370 use_cache_branch = np.array(
371 [True],
372 dtype=bool
373 )
374
375
376
377# ============================================================
378# Decode Output Tokens
379# ============================================================
380
381translation = tokenizer.decode(
382 [
383 token
384 for token in generated_tokens[1:]
385 if token != eos_token_id
386 ],
387 skip_special_tokens=True
388)
389
390
391print("=" * 60)
392print("Translation Result")
393print("=" * 60)
394print(translation)
395print("=" * 60)Helsinki-NLP-opus-mt-en-de with your desired model folder name.Helsinki-NLP-opus-mt-tc-base-bat-zle)