mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-29 20:14:01 +00:00
Attention fusion detect num_heads and hidden_size automatically (#6920)
This commit is contained in:
parent
acfe7ac4ce
commit
4884eee642
8 changed files with 279 additions and 30 deletions
|
|
@ -5,7 +5,8 @@
|
|||
import numpy as np
|
||||
from logging import getLogger
|
||||
from enum import Enum
|
||||
from onnx import helper, numpy_helper, TensorProto
|
||||
from typing import Tuple, Union
|
||||
from onnx import helper, numpy_helper, TensorProto, NodeProto
|
||||
from onnx_model import OnnxModel
|
||||
from fusion_base import Fusion
|
||||
from fusion_utils import FusionUtils
|
||||
|
|
@ -45,7 +46,7 @@ class AttentionMask():
|
|||
assert len(self.mask_indice) > 0
|
||||
return next(iter(self.mask_indice))
|
||||
|
||||
def process_mask(self, input):
|
||||
def process_mask(self, input: str) -> str:
|
||||
if self.mask_format == AttentionMaskFormat.NoMask:
|
||||
return None
|
||||
|
||||
|
|
@ -90,7 +91,62 @@ class FusionAttention(Fusion):
|
|||
self.num_heads = num_heads
|
||||
self.attention_mask = attention_mask
|
||||
|
||||
def create_attention_node(self, mask_index, q_matmul, k_matmul, v_matmul, q_add, k_add, v_add, input, output):
|
||||
def get_num_heads_and_hidden_size(self, reshape_q: NodeProto) -> Tuple[int, int]:
|
||||
""" Detect num_heads and hidden_size from a reshape node.
|
||||
|
||||
Args:
|
||||
reshape_q (NodeProto): reshape node for Q
|
||||
|
||||
Returns:
|
||||
Tuple[int, int]: num_heads and hidden_size
|
||||
"""
|
||||
|
||||
# we assume that reshape fusion has done, so the shape is a tensor like [0, 0, num_heads, head_size]
|
||||
q_shape = self.model.get_initializer(reshape_q.input[1])
|
||||
if q_shape is None:
|
||||
logger.debug(f"{reshape_q.input[1]} is not initializer.")
|
||||
return self.num_heads, self.hidden_size # Fall back to user specified value
|
||||
|
||||
q_shape_value = numpy_helper.to_array(q_shape)
|
||||
if len(q_shape_value) != 4 or (q_shape_value[2] <= 0 or q_shape_value[3] <= 0):
|
||||
logger.debug(f"q_shape_value={q_shape_value}. Expected value are like [0, 0, num_heads, head_size].")
|
||||
return self.num_heads, self.hidden_size # Fall back to user specified value
|
||||
|
||||
num_heads = q_shape_value[2]
|
||||
head_size = q_shape_value[3]
|
||||
hidden_size = num_heads * head_size
|
||||
|
||||
if self.num_heads > 0 and num_heads != self.num_heads:
|
||||
logger.warn("--num_heads is {self.num_heads}. Detected value is {num_heads}. Using detected value.")
|
||||
|
||||
if self.hidden_size > 0 and hidden_size != self.hidden_size:
|
||||
logger.warn("--hidden_size is {self.hidden_size}. Detected value is {hidden_size}. Using detected value.")
|
||||
|
||||
return num_heads, hidden_size
|
||||
|
||||
def create_attention_node(self, mask_index: str, q_matmul: NodeProto, k_matmul: NodeProto, v_matmul: NodeProto,
|
||||
q_add: NodeProto, k_add: NodeProto, v_add: NodeProto, num_heads: int, hidden_size: int,
|
||||
input: str, output: str) -> Union[NodeProto, None]:
|
||||
""" Create an Attention node.
|
||||
|
||||
Args:
|
||||
mask_index (str): mask input
|
||||
q_matmul (NodeProto): MatMul node in fully connection for Q
|
||||
k_matmul (NodeProto): MatMul node in fully connection for K
|
||||
v_matmul (NodeProto): MatMul node in fully connection for V
|
||||
q_add (NodeProto): Add bias node in fully connection for Q
|
||||
k_add (NodeProto): Add bias node in fully connection for K
|
||||
v_add (NodeProto): Add bias node in fully connection for V
|
||||
num_heads (int): number of attention heads. If a model is pruned, it is the number of heads after pruning.
|
||||
hidden_size (int): hidden dimension. If a model is pruned, it is the hidden dimension after pruning.
|
||||
input (str): input name
|
||||
output (str): output name
|
||||
|
||||
Returns:
|
||||
Union[NodeProto, None]: the node created or None if failed.
|
||||
"""
|
||||
assert num_heads > 0 or hidden_size > 0 or (hidden_size % num_heads) == 0
|
||||
|
||||
q_weight = self.model.get_initializer(q_matmul.input[1])
|
||||
k_weight = self.model.get_initializer(k_matmul.input[1])
|
||||
v_weight = self.model.get_initializer(v_matmul.input[1])
|
||||
|
|
@ -103,10 +159,6 @@ class FusionAttention(Fusion):
|
|||
return None
|
||||
if not (k_weight and v_weight and q_bias and k_bias):
|
||||
return None
|
||||
|
||||
assert (self.hidden_size % self.num_heads) == 0
|
||||
head_hidden_size = self.hidden_size // self.num_heads
|
||||
|
||||
qw = numpy_helper.to_array(q_weight)
|
||||
kw = numpy_helper.to_array(k_weight)
|
||||
vw = numpy_helper.to_array(v_weight)
|
||||
|
|
@ -114,8 +166,18 @@ class FusionAttention(Fusion):
|
|||
# Check if all matrices have the same shape
|
||||
assert qw.shape == kw.shape == vw.shape
|
||||
|
||||
if len(qw.shape) != 2:
|
||||
logger.debug(f"weights for Q is expected to be 2D.")
|
||||
return None
|
||||
|
||||
# All the matrices have the same shape (in_size, out_size)
|
||||
in_size, out_size = qw.shape
|
||||
|
||||
if out_size != hidden_size:
|
||||
logger.debug(
|
||||
f"Shape for weights of Q is {in_size, out_size}, which does not match hidden_size={hidden_size}")
|
||||
return None
|
||||
|
||||
qkv_weight = np.stack((qw, kw, vw), axis=-2)
|
||||
|
||||
qb = numpy_helper.to_array(q_bias)
|
||||
|
|
@ -135,6 +197,7 @@ class FusionAttention(Fusion):
|
|||
data_type=TensorProto.FLOAT,
|
||||
dims=[in_size, 3 * out_size],
|
||||
vals=qkv_weight.flatten().tolist())
|
||||
|
||||
# Sometimes weights and bias are stored in fp16
|
||||
if q_weight.data_type == 10:
|
||||
weight.CopyFrom(numpy_helper.from_array(numpy_helper.to_array(weight).astype(np.float16), weight.name))
|
||||
|
|
@ -157,7 +220,7 @@ class FusionAttention(Fusion):
|
|||
outputs=[output],
|
||||
name=attention_node_name)
|
||||
attention_node.domain = "com.microsoft"
|
||||
attention_node.attribute.extend([helper.make_attribute("num_heads", out_size // head_hidden_size)])
|
||||
attention_node.attribute.extend([helper.make_attribute("num_heads", num_heads)])
|
||||
|
||||
return attention_node
|
||||
|
||||
|
|
@ -257,6 +320,7 @@ class FusionAttention(Fusion):
|
|||
if q_nodes is None:
|
||||
logger.debug("fuse_attention: failed to match q path")
|
||||
return
|
||||
reshape_q = q_nodes[-3]
|
||||
add_q = q_nodes[-2]
|
||||
matmul_q = q_nodes[-1]
|
||||
|
||||
|
|
@ -290,8 +354,13 @@ class FusionAttention(Fusion):
|
|||
|
||||
attention_last_node = reshape_qkv if einsum_node is None else transpose_qkv
|
||||
|
||||
num_heads, hidden_size = self.get_num_heads_and_hidden_size(reshape_q)
|
||||
if num_heads <= 0 or hidden_size <= 0 or (hidden_size % num_heads) != 0:
|
||||
logger.debug("fuse_attention: failed to detect num_heads or hidden_size")
|
||||
return
|
||||
|
||||
new_node = self.create_attention_node(mask_index, matmul_q, matmul_k, matmul_v, add_q, add_k, add_v,
|
||||
root_input, attention_last_node.output[0])
|
||||
num_heads, hidden_size, root_input, attention_last_node.output[0])
|
||||
if new_node is None:
|
||||
return
|
||||
|
||||
|
|
@ -303,9 +372,8 @@ class FusionAttention(Fusion):
|
|||
shape_tensor = helper.make_tensor(name="shape_modified_tensor" + unique_index,
|
||||
data_type=TensorProto.INT64,
|
||||
dims=[4],
|
||||
vals=np.int64(
|
||||
[0, 0, self.num_heads,
|
||||
int(self.hidden_size / self.num_heads)]).tobytes(),
|
||||
vals=np.int64([0, 0, num_heads,
|
||||
int(hidden_size / num_heads)]).tobytes(),
|
||||
raw=True)
|
||||
self.model.add_initializer(shape_tensor)
|
||||
self.model.add_node(
|
||||
|
|
|
|||
|
|
@ -19,6 +19,7 @@ class FusionGptAttention(Fusion):
|
|||
"""
|
||||
def __init__(self, model: OnnxModel, num_heads: int):
|
||||
super().__init__(model, "Attention", "LayerNormalization", "with past")
|
||||
# TODO: detect num_heads from graph like FusionAttention
|
||||
self.num_heads = num_heads
|
||||
self.utils = FusionUtils(model)
|
||||
self.casted_attention_mask = {} # map from name of attention mask to the name that casted to int32
|
||||
|
|
|
|||
|
|
@ -19,6 +19,7 @@ class FusionGptAttentionNoPast(Fusion):
|
|||
"""
|
||||
def __init__(self, model: OnnxModel, num_heads: int):
|
||||
super().__init__(model, "Attention", "LayerNormalization", "without past")
|
||||
# TODO: detect num_heads from graph like FusionAttention
|
||||
self.num_heads = num_heads
|
||||
|
||||
def create_attention_node(self, gemm, gemm_qkv, input, output):
|
||||
|
|
|
|||
|
|
@ -5,7 +5,7 @@
|
|||
|
||||
from logging import getLogger
|
||||
from typing import List
|
||||
from onnx import TensorProto, helper
|
||||
from onnx import ModelProto, TensorProto, helper
|
||||
from onnx_model import OnnxModel
|
||||
from fusion_reshape import FusionReshape
|
||||
from fusion_layernorm import FusionLayerNormalization, FusionLayerNormalizationTF
|
||||
|
|
@ -47,9 +47,15 @@ class BertOptimizationOptions:
|
|||
|
||||
|
||||
class BertOnnxModel(OnnxModel):
|
||||
def __init__(self, model, num_heads, hidden_size):
|
||||
assert num_heads > 0
|
||||
assert hidden_size % num_heads == 0
|
||||
def __init__(self, model: ModelProto, num_heads: int = 0, hidden_size: int = 0):
|
||||
"""Initialize BERT ONNX Model.
|
||||
|
||||
Args:
|
||||
model (ModelProto): the ONNX model
|
||||
num_heads (int, optional): number of attentioin heads. Defaults to 0, and we will detect the parameter automatically.
|
||||
hidden_size (int, optional): hidden dimension. Defaults to 0, and we will detect the parameter automatically.
|
||||
"""
|
||||
assert (num_heads == 0 and hidden_size == 0) or (num_heads > 0 and hidden_size % num_heads == 0)
|
||||
|
||||
super().__init__(model)
|
||||
self.num_heads = num_heads
|
||||
|
|
|
|||
|
|
@ -119,17 +119,23 @@ def _parse_arguments():
|
|||
choices=list(MODEL_CLASSES.keys()),
|
||||
help="Model type selected in the list: " + ", ".join(MODEL_CLASSES.keys()))
|
||||
|
||||
parser.add_argument('--num_heads',
|
||||
required=False,
|
||||
type=int,
|
||||
default=12,
|
||||
help="number of attention heads. 12 for bert-base model and 16 for bert-large")
|
||||
parser.add_argument(
|
||||
'--num_heads',
|
||||
required=False,
|
||||
type=int,
|
||||
default=12,
|
||||
help=
|
||||
"number of attention heads. 12 for bert-base model and 16 for bert-large. For BERT, set it to 0 to detect automatically."
|
||||
)
|
||||
|
||||
parser.add_argument('--hidden_size',
|
||||
required=False,
|
||||
type=int,
|
||||
default=768,
|
||||
help="bert model hidden size. 768 for bert-base model and 1024 for bert-large")
|
||||
parser.add_argument(
|
||||
'--hidden_size',
|
||||
required=False,
|
||||
type=int,
|
||||
default=768,
|
||||
help=
|
||||
"bert model hidden size. 768 for bert-base model and 1024 for bert-large. For BERT, set it to 0 to detect automatically."
|
||||
)
|
||||
|
||||
parser.add_argument('--input_int32',
|
||||
required=False,
|
||||
|
|
@ -253,8 +259,8 @@ def _get_optimization_options(args):
|
|||
|
||||
def optimize_model(input,
|
||||
model_type='bert',
|
||||
num_heads=12,
|
||||
hidden_size=768,
|
||||
num_heads=0,
|
||||
hidden_size=0,
|
||||
optimization_options=None,
|
||||
opt_level=0,
|
||||
use_gpu=False,
|
||||
|
|
@ -269,8 +275,8 @@ def optimize_model(input,
|
|||
Args:
|
||||
input (str): input model path.
|
||||
model_type (str): model type - like bert, bert_tf, bert_keras or gpt2.
|
||||
num_heads (int): number of attention heads.
|
||||
hidden_size (int): hidden size.
|
||||
num_heads (int): number of attention heads. Default is 0 to allow detect the parameter from graph automatically (for model_type "bert" only).
|
||||
hidden_size (int): hidden size. Default is 0 to allow detect the parameter from graph automatically (for model_type "bert" only).
|
||||
optimization_options (OptimizationOptions or None): optimization options that can use to turn on/off some fusions.
|
||||
opt_level (int): onnxruntime graph optimization level (0, 1, 2 or 99). When the level > 0, onnxruntime will be used to optimize model first.
|
||||
use_gpu (bool): use gpu or not for onnxruntime.
|
||||
|
|
|
|||
|
|
@ -0,0 +1,133 @@
|
|||
# -------------------------------------------------------------------------
|
||||
# Copyright (c) Microsoft Corporation. All rights reserved.
|
||||
# Licensed under the MIT License. See License.txt in the project root for
|
||||
# license information.
|
||||
# --------------------------------------------------------------------------
|
||||
|
||||
import onnx
|
||||
import math
|
||||
from typing import List
|
||||
from packaging import version
|
||||
from onnx import helper, TensorProto
|
||||
|
||||
|
||||
def float_tensor(name: str, shape: List[int], random=False):
|
||||
low = 0.0
|
||||
high = 1.0
|
||||
total_elements = 1
|
||||
for x in shape:
|
||||
total_elements *= x
|
||||
weights = [random.uniform(low, high) for _ in range(total_elements)] if random else [1.0] * total_elements
|
||||
return helper.make_tensor(name, TensorProto.FLOAT, shape, weights)
|
||||
|
||||
|
||||
def create_bert_attention(input_hidden_size=16, pruned_num_heads=2, pruned_head_size=4, use_float_mask=False):
|
||||
# unsqueeze in opset version 13 has two inputs (axis is moved from attribute to input).
|
||||
has_unsqueeze_two_inputs = (version.parse(onnx.__version__) >= version.parse('1.8.0'))
|
||||
|
||||
# nodes in attention subgraph
|
||||
nodes = [
|
||||
helper.make_node("Add", ["input_1", "input_2"], ["layernorm_input"], "add_layernorm"),
|
||||
helper.make_node("LayerNormalization", ["layernorm_input", "layer_norm_weight", "layer_norm_bias"],
|
||||
["layernorm_out"],
|
||||
"layernorm",
|
||||
axis=-1,
|
||||
epsion=0.000009999999747378752),
|
||||
|
||||
# q nodes
|
||||
helper.make_node("MatMul", ["layernorm_out", "matmul_q_weight"], ["matmul_q_out"], "matmul_q"),
|
||||
helper.make_node("Add", ["matmul_q_out", "add_q_weight"], ["add_q_out"], "add_q"),
|
||||
helper.make_node("Reshape", ["add_q_out", "reshape_weight_1"], ["reshape_q_out"], "reshape_q"),
|
||||
helper.make_node("Transpose", ["reshape_q_out"], ["transpose_q_out"], "transpose_q", perm=[0, 2, 1, 3]),
|
||||
|
||||
# k nodes
|
||||
helper.make_node("MatMul", ["layernorm_out", "matmul_k_weight"], ["matmul_k_out"], "matmul_k"),
|
||||
helper.make_node("Add", ["matmul_k_out", "add_k_weight"], ["add_k_out"], "add_k"),
|
||||
helper.make_node("Reshape", ["add_k_out", "reshape_weight_1"], ["reshape_k_out"], "reshape_k"),
|
||||
helper.make_node("Transpose", ["reshape_k_out"], ["transpose_k_out"], "transpose_k", perm=[0, 2, 3, 1]),
|
||||
|
||||
# mask nodes
|
||||
helper.make_node("Unsqueeze", ["input_mask", "axes_1"], ["unsqueeze0_out"], "unsqueeze0") if has_unsqueeze_two_inputs \
|
||||
else helper.make_node("Unsqueeze", ["input_mask"], ["unsqueeze0_out"], "unsqueeze0", axes=[1]),
|
||||
helper.make_node("Unsqueeze", ["unsqueeze0_out", "axes_2"], ["unsqueeze1_out"], "unsqueeze1") if has_unsqueeze_two_inputs \
|
||||
else helper.make_node("Unsqueeze", ["unsqueeze0_out"], ["unsqueeze1_out"], "unsqueeze1", axes=[2]),
|
||||
|
||||
# when attention_mask is float type, no need to cast
|
||||
helper.make_node("Cast", ["unsqueeze1_out"], ["cast_out"], "cast", to=1) if not use_float_mask else None,
|
||||
helper.make_node("Sub", ["sub_weight", "unsqueeze1_out" if use_float_mask else "cast_out"], ["sub_out"], "sub"),
|
||||
helper.make_node("Mul", ["sub_out", "mul_weight"], ["mul_mask_out"], "mul_mask"),
|
||||
|
||||
# qk nodes
|
||||
helper.make_node("MatMul", ["transpose_q_out", "transpose_k_out"], ["matmul_qk_out"], "matmul_qk"),
|
||||
helper.make_node("Div", ["matmul_qk_out", "div_weight"], ["div_qk_out"], "div_qk"),
|
||||
helper.make_node("Add", ["div_qk_out", "mul_mask_out"], ["add_qk_out"], "add_qk"),
|
||||
helper.make_node("Softmax", ["add_qk_out"], ["softmax_qk_out"], "softmax_qk", axis=3),
|
||||
|
||||
# v nodes
|
||||
helper.make_node("MatMul", ["layernorm_out", "matmul_v_weight"], ["matmul_v_out"], "matmul_v"),
|
||||
helper.make_node("Add", ["matmul_v_out", "add_v_weight"], ["add_v_out"], "add_v"),
|
||||
helper.make_node("Reshape", ["add_v_out", "reshape_weight_1"], ["reshape_v_out"], "reshape_v"),
|
||||
helper.make_node("Transpose", ["reshape_v_out"], ["transpose_v_out"], "transpose_v", perm=[0, 2, 1, 3]),
|
||||
|
||||
# qkv nodes
|
||||
helper.make_node("MatMul", ["softmax_qk_out", "transpose_v_out"], ["matmul_qkv_1_out"], "matmul_qkv_1"),
|
||||
helper.make_node("Transpose", ["matmul_qkv_1_out"], ["transpose_qkv_out"], "transpose_qkv", perm=[0, 2, 1, 3]),
|
||||
helper.make_node("Reshape", ["transpose_qkv_out", "reshape_weight_2"], ["reshape_qkv_out"], "reshape_qkv"),
|
||||
helper.make_node("MatMul", ["reshape_qkv_out", "matmul_qkv_weight"], ["matmul_qkv_2_out"], "matmul_qkv_2"),
|
||||
helper.make_node("Add", ["matmul_qkv_2_out", "add_qkv_weight"], ["add_qkv_out"], "add_qkv"),
|
||||
helper.make_node("Add", ["add_qkv_out", "layernorm_out"], ["skip_output"], "add_skip"),
|
||||
helper.make_node("LayerNormalization", ["skip_output", "layer_norm_weight", "layer_norm_bias"], ["output"],
|
||||
"layernorm2",
|
||||
axis=-1,
|
||||
epsion=0.000009999999747378752),
|
||||
]
|
||||
|
||||
pruned_hidden_size = pruned_num_heads * pruned_head_size
|
||||
initializers = [ # initializers
|
||||
float_tensor('layer_norm_weight', [input_hidden_size]),
|
||||
float_tensor('layer_norm_bias', [input_hidden_size]),
|
||||
float_tensor('matmul_q_weight', [input_hidden_size, pruned_hidden_size]),
|
||||
float_tensor('matmul_k_weight', [input_hidden_size, pruned_hidden_size]),
|
||||
float_tensor('matmul_v_weight', [input_hidden_size, pruned_hidden_size]),
|
||||
float_tensor('matmul_qkv_weight', [pruned_hidden_size, input_hidden_size]),
|
||||
float_tensor('add_q_weight', [pruned_hidden_size]),
|
||||
float_tensor('add_k_weight', [pruned_hidden_size]),
|
||||
float_tensor('add_v_weight', [pruned_hidden_size]),
|
||||
float_tensor('add_qkv_weight', [input_hidden_size]),
|
||||
helper.make_tensor('div_weight', TensorProto.FLOAT, [1], [math.sqrt(pruned_head_size)]),
|
||||
helper.make_tensor('sub_weight', TensorProto.FLOAT, [1], [1.0]),
|
||||
helper.make_tensor('mul_weight', TensorProto.FLOAT, [1], [-10000]),
|
||||
helper.make_tensor('reshape_weight_1', TensorProto.INT64, [4], [0, 0, pruned_num_heads, pruned_head_size]),
|
||||
helper.make_tensor('reshape_weight_2', TensorProto.INT64, [3], [0, 0, pruned_hidden_size]),
|
||||
]
|
||||
|
||||
if has_unsqueeze_two_inputs:
|
||||
initializers.append(helper.make_tensor('axes_1', TensorProto.INT64, [1], [1]))
|
||||
initializers.append(helper.make_tensor('axes_2', TensorProto.INT64, [1], [2]))
|
||||
|
||||
batch_size = 1
|
||||
sequence_length = 3
|
||||
graph = helper.make_graph(
|
||||
[node for node in nodes if node],
|
||||
"AttentionFusionPrunedModel", #name
|
||||
[ # inputs
|
||||
helper.make_tensor_value_info('input_1', TensorProto.FLOAT,
|
||||
[batch_size, sequence_length, input_hidden_size]),
|
||||
helper.make_tensor_value_info('input_2', TensorProto.FLOAT,
|
||||
[batch_size, sequence_length, input_hidden_size]),
|
||||
helper.make_tensor_value_info('input_mask', TensorProto.FLOAT if use_float_mask else TensorProto.INT64,
|
||||
[batch_size, sequence_length])
|
||||
],
|
||||
[ # outputs
|
||||
helper.make_tensor_value_info('output', TensorProto.FLOAT,
|
||||
[batch_size, sequence_length, input_hidden_size]),
|
||||
],
|
||||
initializers)
|
||||
|
||||
model = helper.make_model(graph)
|
||||
return model
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
model = create_bert_attention()
|
||||
onnx.save(model, "pruned_bert_attention.onnx")
|
||||
|
|
@ -0,0 +1,34 @@
|
|||
# -------------------------------------------------------------------------
|
||||
# Copyright (c) Microsoft Corporation. All rights reserved.
|
||||
# Licensed under the MIT License. See License.txt in the project root for
|
||||
# license information.
|
||||
# --------------------------------------------------------------------------
|
||||
|
||||
import unittest
|
||||
import os
|
||||
import sys
|
||||
import onnx
|
||||
from bert_model_generator import create_bert_attention
|
||||
|
||||
# set path so that we could import from parent directory
|
||||
sys.path.append(os.path.join(os.path.dirname(__file__), '..'))
|
||||
from optimizer import optimize_model
|
||||
|
||||
|
||||
class TestFusion(unittest.TestCase):
|
||||
def test_attention_fusion_pruned_model(self):
|
||||
model = create_bert_attention()
|
||||
dir = '.'
|
||||
model_path = os.path.join(dir, "pruned_attention.onnx")
|
||||
onnx.save(model, model_path)
|
||||
optimized_model = optimize_model(model_path)
|
||||
os.remove(model_path)
|
||||
|
||||
expected_model_path = os.path.join(os.path.dirname(__file__), 'test_data', 'fusion',
|
||||
'pruned_attention_opt.onnx')
|
||||
expected = onnx.load(expected_model_path)
|
||||
self.assertEqual(str(optimized_model.model.graph), str(expected.graph))
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
unittest.main()
|
||||
Binary file not shown.
Loading…
Reference in a new issue