Views
No views yet
distilbert/distilgpt2
that includes attention weights as graph outputs, quantized to int8 for efficient
in-browser inference via Transformers.js.attentions tuple
returned by model(..., output_attentions=True)) because they aren't needed for
inference. This export preserves them as named graph outputs so the model can be used
for educational visualizations of how attention works in causal language models.distilbert/distilgpt2 (82M parameters)| Name | Shape | Description |
|---|---|---|
input_ids | [batch, seq] | Token IDs |
attention_mask | [batch, seq] | All ones for non-padded inputs |
position_ids | [batch, seq] | [0, 1, ..., seq-1] |
logits | [batch, seq, 50257] | Next-token logits |
attentions.0 ... attentions.5 | [batch, 12, seq, seq] | Attention weights per layer |
1import { AutoTokenizer, AutoModel, Tensor } from "@huggingface/transformers";
2
3const tokenizer = await AutoTokenizer.from_pretrained("dbernsohn/distilgpt2-onnx");
4const model = await AutoModel.from_pretrained("dbernsohn/distilgpt2-onnx", {
5 dtype: "fp32",
6 model_file_name: "model",
7});
8
9const inputs = tokenizer("The capital of France is");
10const seqLen = inputs.input_ids.dims[1];
11
12const result = await model.forward({
13 input_ids: inputs.input_ids,
14 attention_mask: new Tensor(
15 "int64",
16 new BigInt64Array(seqLen).fill(1n),
17 [1, seqLen]
18 ),
19 position_ids: new Tensor(
20 "int64",
21 BigInt64Array.from({ length: seqLen }, (_, i) => BigInt(i)),
22 [1, seqLen]
23 ),
24});
25
26// result.logits: [1, seq, 50257]
27// result["attentions.0"] ... result["attentions.5"]: [1, 12, seq, seq]1import numpy as np
2import onnxruntime as ort
3from transformers import AutoTokenizer
4
5tokenizer = AutoTokenizer.from_pretrained("dbernsohn/distilgpt2-onnx")
6session = ort.InferenceSession("onnx/model.onnx")
7
8ids = tokenizer("The capital of France is", return_tensors="np")
9seq_len = ids["input_ids"].shape[1]
10
11outputs = session.run(None, {
12 "input_ids": ids["input_ids"].astype(np.int64),
13 "attention_mask": np.ones((1, seq_len), dtype=np.int64),
14 "position_ids": np.arange(seq_len, dtype=np.int64).reshape(1, -1),
15})
16
17logits = outputs[0] # [1, seq, 50257]
18attentions = outputs[1:] # 6 tensors of [1, 12, seq, seq]pip install torch transformers optimum[onnxruntime] onnx onnxscript1import shutil
2from pathlib import Path
3import torch
4import onnx
5from onnxruntime.quantization import quantize_dynamic, QuantType
6from transformers import AutoTokenizer, AutoModelForCausalLM
7
8MODEL_ID = "distilbert/distilgpt2"
9OUTPUT_DIR = Path("exported-distilgpt2")
10
11if OUTPUT_DIR.exists():
12 shutil.rmtree(OUTPUT_DIR)
13OUTPUT_DIR.mkdir()
14(OUTPUT_DIR / "onnx").mkdir()
15
16tokenizer = AutoTokenizer.from_pretrained(MODEL_ID)
17model = AutoModelForCausalLM.from_pretrained(MODEL_ID, attn_implementation="eager")
18model.eval()
19
20# Wrapper that returns logits + attention layers as separate outputs
21class AttentionWrapper(torch.nn.Module):
22 def __init__(self, base_model):
23 super().__init__()
24 self.base = base_model
25
26 def forward(self, input_ids, attention_mask, position_ids):
27 out = self.base(
28 input_ids=input_ids,
29 attention_mask=attention_mask,
30 position_ids=position_ids,
31 output_attentions=True,
32 use_cache=False,
33 )
34 return (out.logits,) + out.attentions
35
36wrapper = AttentionWrapper(model)
37wrapper.eval()
38
39dummy = tokenizer("Hello world", return_tensors="pt")
40input_ids = dummy["input_ids"]
41attention_mask = dummy["attention_mask"]
42seq_len = input_ids.shape[1]
43position_ids = torch.arange(seq_len, dtype=torch.long).unsqueeze(0)
44
45with torch.no_grad():
46 test_out = wrapper(input_ids, attention_mask, position_ids)
47num_attn_layers = len(test_out) - 1
48
49output_names = ["logits"]
50dynamic_axes = {
51 "input_ids": {0: "batch", 1: "seq"},
52 "attention_mask": {0: "batch", 1: "seq"},
53 "position_ids": {0: "batch", 1: "seq"},
54 "logits": {0: "batch", 1: "seq"},
55}
56for i in range(num_attn_layers):
57 name = f"attentions.{i}"
58 output_names.append(name)
59 dynamic_axes[name] = {0: "batch", 2: "seq", 3: "seq"}
60
61onnx_path = OUTPUT_DIR / "onnx" / "model_fp32.onnx"
62
63torch.onnx.export(
64 wrapper,
65 (input_ids, attention_mask, position_ids),
66 str(onnx_path),
67 input_names=["input_ids", "attention_mask", "position_ids"],
68 output_names=output_names,
69 dynamic_axes=dynamic_axes,
70 opset_version=17,
71 do_constant_folding=True,
72 dynamo=True,
73 optimize=False, # skip optimization that fails on GPT-2 + attention outputs
74)
75
76# Merge external data into single file
77data_file = onnx_path.parent / (onnx_path.name + ".data")
78if data_file.exists():
79 m = onnx.load(str(onnx_path), load_external_data=True)
80 onnx.save_model(m, str(onnx_path), save_as_external_data=False)
81 data_file.unlink()
82
83# Quantize to int8
84quantized_path = OUTPUT_DIR / "onnx" / "model.onnx"
85quantize_dynamic(
86 str(onnx_path),
87 str(quantized_path),
88 weight_type=QuantType.QInt8,
89)
90onnx_path.unlink()
91
92# Save tokenizer + config
93tokenizer.save_pretrained(str(OUTPUT_DIR))
94model.config.save_pretrained(str(OUTPUT_DIR))