mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-30 20:18:08 +00:00
Allow using separate GPT2 decoder subgraphs for the initial run and the subsequent runs in BeamSearch/GreedySearch (#13914)
This commit is contained in:
parent
22fa62152a
commit
51aaf2e021
2 changed files with 286 additions and 8 deletions
|
|
@ -182,12 +182,21 @@ def parse_arguments(argv: Optional[List[str]] = None) -> argparse.Namespace:
|
|||
)
|
||||
output_group.set_defaults(pad_vocab_size=True)
|
||||
|
||||
output_group.add_argument(
|
||||
"-sgd",
|
||||
"--separate_gpt2_decoder_for_init_run",
|
||||
required=False,
|
||||
action="store_true",
|
||||
help="Have separate decoder subgraphs for initial and remaining runs. This allows for optimizations based on sequence lengths in each subgraph",
|
||||
)
|
||||
output_group.set_defaults(separate_gpt2_decoder_for_init_run=False)
|
||||
|
||||
output_group.add_argument(
|
||||
"-i",
|
||||
"--disable_shared_initializers",
|
||||
required=False,
|
||||
action="store_true",
|
||||
help="do not share initializers in encoder and decoder. It will increase memory usage of t5/mt5 models.",
|
||||
help="do not share initializers in encoder and decoder for T5 or in the init decoder and decoder for GPT2. It will increase memory usage of t5/mt5/gpt2 models.",
|
||||
)
|
||||
output_group.set_defaults(disable_shared_initializers=False)
|
||||
|
||||
|
|
@ -867,6 +876,218 @@ def move_initializers(
|
|||
return moved_initializers
|
||||
|
||||
|
||||
def update_input_shapes_for_gpt2_decoder_model(decoder_onnx_path: str, use_external_data_format: bool = True):
|
||||
"""Update the input shapes for the inputs "input_ids" and "position_ids" and make the sequence length dim value 1 for each of them.
|
||||
The decoder model will be over-written.
|
||||
|
||||
Args:
|
||||
decoder_onnx_path (str): Path of GPT-2 decoder onnx model
|
||||
use_external_data_format(bool): output tensors to external data or not.
|
||||
"""
|
||||
|
||||
decoder_model_proto = onnx.load_model(decoder_onnx_path, load_external_data=True)
|
||||
for i in range(len(decoder_model_proto.graph.input)):
|
||||
if (
|
||||
decoder_model_proto.graph.input[i].name == "input_ids"
|
||||
or decoder_model_proto.graph.input[i].name == "position_ids"
|
||||
):
|
||||
shape_dim_proto = decoder_model_proto.graph.input[i].type.tensor_type.shape.dim[1]
|
||||
|
||||
# Clear any existing dim_param first
|
||||
if shape_dim_proto.HasField("dim_param"):
|
||||
shape_dim_proto.Clear()
|
||||
|
||||
# Update dim_value to be 1
|
||||
shape_dim_proto.dim_value = 1
|
||||
|
||||
OnnxModel.save(decoder_model_proto, decoder_onnx_path, save_as_external_data=use_external_data_format)
|
||||
return True
|
||||
|
||||
|
||||
def generate_gpt2_init_decoder(
|
||||
decoder_onnx_path: str, init_decoder_onnx_path: str, use_external_data_format: bool = True
|
||||
) -> bool:
|
||||
"""Generates the initial decoder GPT2 subgraph and saves it for downstream use.
|
||||
The initial decoder model will be saved to init_decoder_onnx_path.
|
||||
|
||||
Args:
|
||||
decoder_onnx_path (str): Path of GPT-2 decoder onnx model
|
||||
init_decoder_onnx_path (str): Path of GPT-2 init decoder onnx model
|
||||
use_external_data_format(bool): output tensors to external data or not.
|
||||
"""
|
||||
init_decoder_model_proto = onnx.load_model(decoder_onnx_path, load_external_data=True)
|
||||
|
||||
logits_output_name = init_decoder_model_proto.graph.output[0].name
|
||||
|
||||
gpt2_init_decoder_model = OnnxModel(init_decoder_model_proto)
|
||||
|
||||
output_name_to_node = gpt2_init_decoder_model.output_name_to_node()
|
||||
assert logits_output_name in output_name_to_node
|
||||
|
||||
logits_matmul_node = output_name_to_node[logits_output_name]
|
||||
|
||||
# Sanity check - the logits need to be produced by a MatMul node
|
||||
if logits_matmul_node.op_type != "MatMul":
|
||||
return False
|
||||
|
||||
# Try to find the last residual Add
|
||||
# For fp16, there are Casts along the way
|
||||
logits_matmul_to_residual_add_path = gpt2_init_decoder_model.match_parent_path(
|
||||
logits_matmul_node,
|
||||
[
|
||||
"Cast",
|
||||
"LayerNormalization",
|
||||
"Add",
|
||||
"Add",
|
||||
"Cast",
|
||||
"MatMul",
|
||||
"Cast",
|
||||
"FastGelu",
|
||||
"Cast",
|
||||
"MatMul",
|
||||
"Cast",
|
||||
"LayerNormalization",
|
||||
"Add",
|
||||
],
|
||||
[0, 0, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0],
|
||||
)
|
||||
|
||||
# Try without the Casts before and after the MatMuls
|
||||
if logits_matmul_to_residual_add_path is None:
|
||||
logits_matmul_to_residual_add_path = gpt2_init_decoder_model.match_parent_path(
|
||||
logits_matmul_node,
|
||||
["LayerNormalization", "Add", "Add", "MatMul", "FastGelu", "MatMul", "LayerNormalization", "Add"],
|
||||
[0, 0, 1, 0, 0, 0, 0, 0],
|
||||
)
|
||||
|
||||
# TODO(hasesh): Are there more permutations to try before returning ?
|
||||
if logits_matmul_to_residual_add_path is None:
|
||||
return False
|
||||
|
||||
residual_add_node = logits_matmul_to_residual_add_path[-1]
|
||||
|
||||
residual_add_to_attention_parent_index = 0
|
||||
residual_add_to_attention_path = gpt2_init_decoder_model.match_parent_path(
|
||||
residual_add_node, ["Add", "Cast", "MatMul", "Attention"], [residual_add_to_attention_parent_index, 0, 0, 0]
|
||||
)
|
||||
|
||||
# Try other parent index of the residual Add node
|
||||
if residual_add_to_attention_path is None:
|
||||
residual_add_to_attention_parent_index = 1
|
||||
residual_add_to_attention_path = gpt2_init_decoder_model.match_parent_path(
|
||||
residual_add_node, ["Add", "Cast", "MatMul", "Attention"], [residual_add_to_attention_parent_index, 0, 0, 0]
|
||||
)
|
||||
|
||||
# Try without the Casts before and after the MatMuls
|
||||
if residual_add_to_attention_path is None:
|
||||
residual_add_to_attention_parent_index = 0
|
||||
residual_add_to_attention_path = gpt2_init_decoder_model.match_parent_path(
|
||||
residual_add_node, ["Add", "MatMul", "Attention"], [residual_add_to_attention_parent_index, 0, 0]
|
||||
)
|
||||
|
||||
# Try without the Casts before and after the MatMuls and other parent index of the residual Add node
|
||||
if residual_add_to_attention_path is None:
|
||||
residual_add_to_attention_parent_index = 1
|
||||
residual_add_to_attention_path = gpt2_init_decoder_model.match_parent_path(
|
||||
residual_add_node, ["Add", "MatMul", "Attention"], [residual_add_to_attention_parent_index, 0, 0]
|
||||
)
|
||||
|
||||
# TODO(hasesh): Are there more permutations to try before returning ?
|
||||
if residual_add_to_attention_path is None:
|
||||
return False
|
||||
|
||||
residual_add_to_add_parent_index = 0 if residual_add_to_attention_parent_index == 1 else 1
|
||||
add_before_residual_add = gpt2_init_decoder_model.match_parent(
|
||||
residual_add_node, "Add", residual_add_to_add_parent_index
|
||||
)
|
||||
|
||||
if add_before_residual_add is None:
|
||||
return False
|
||||
|
||||
attention = residual_add_to_attention_path[-1]
|
||||
matmul_after_attention = residual_add_to_attention_path[-2]
|
||||
|
||||
slice_starts = onnx.helper.make_tensor(
|
||||
name="SliceLastTokenStarts",
|
||||
data_type=TensorProto.INT32,
|
||||
dims=[1],
|
||||
vals=[-1],
|
||||
)
|
||||
|
||||
slice_ends = onnx.helper.make_tensor(
|
||||
name="SliceLastTokenEnds",
|
||||
data_type=TensorProto.INT32,
|
||||
dims=[1],
|
||||
vals=[-2],
|
||||
)
|
||||
|
||||
slice_axes = onnx.helper.make_tensor(
|
||||
name="SliceLastTokenAxes",
|
||||
data_type=TensorProto.INT32,
|
||||
dims=[1],
|
||||
vals=[1],
|
||||
)
|
||||
|
||||
slice_steps = onnx.helper.make_tensor(
|
||||
name="SliceLastTokenSteps",
|
||||
data_type=TensorProto.INT32,
|
||||
dims=[1],
|
||||
vals=[-1],
|
||||
)
|
||||
|
||||
gpt2_init_decoder_model.add_initializer(slice_starts)
|
||||
gpt2_init_decoder_model.add_initializer(slice_ends)
|
||||
gpt2_init_decoder_model.add_initializer(slice_axes)
|
||||
gpt2_init_decoder_model.add_initializer(slice_steps)
|
||||
|
||||
# Add Slice node to the graph such that it consumes the output of Attention
|
||||
slice_0_output_name = "edge_modified_" + attention.output[0]
|
||||
slice_node_0 = onnx.helper.make_node(
|
||||
"Slice",
|
||||
inputs=[
|
||||
attention.output[0],
|
||||
"SliceLastTokenStarts",
|
||||
"SliceLastTokenEnds",
|
||||
"SliceLastTokenAxes",
|
||||
"SliceLastTokenSteps",
|
||||
],
|
||||
outputs=[slice_0_output_name],
|
||||
name=gpt2_init_decoder_model.create_node_name("Slice", "GatherLastToken_0_"),
|
||||
)
|
||||
|
||||
# Add Slice node to the graph such that it consumes the output of Add before the residual Add
|
||||
slice_1_output_name = "edge_modified_" + add_before_residual_add.output[0]
|
||||
slice_node_1 = onnx.helper.make_node(
|
||||
"Slice",
|
||||
inputs=[
|
||||
add_before_residual_add.output[0],
|
||||
"SliceLastTokenStarts",
|
||||
"SliceLastTokenEnds",
|
||||
"SliceLastTokenAxes",
|
||||
"SliceLastTokenSteps",
|
||||
],
|
||||
outputs=[slice_1_output_name],
|
||||
name=gpt2_init_decoder_model.create_node_name("Slice", "GatherLastToken_1_"),
|
||||
)
|
||||
|
||||
# Add the 2 Slice nodes
|
||||
gpt2_init_decoder_model.add_node(slice_node_0)
|
||||
gpt2_init_decoder_model.add_node(slice_node_1)
|
||||
|
||||
# Adjust the input(s) to the nodes consuming the outputs of the added Slice nodes
|
||||
gpt2_init_decoder_model.replace_node_input(matmul_after_attention, attention.output[0], slice_0_output_name)
|
||||
gpt2_init_decoder_model.replace_node_input(
|
||||
residual_add_node, add_before_residual_add.output[0], slice_1_output_name
|
||||
)
|
||||
|
||||
# Topologically sort the updated graph
|
||||
gpt2_init_decoder_model.topological_sort()
|
||||
|
||||
# Save the init decoder model
|
||||
OnnxModel.save(init_decoder_model_proto, init_decoder_onnx_path, save_as_external_data=use_external_data_format)
|
||||
return True
|
||||
|
||||
|
||||
def convert_generation_model(args: argparse.Namespace, generation_type: GenerationType = GenerationType.BEAMSEARCH):
|
||||
"""Convert model according to command line arguments.
|
||||
|
||||
|
|
@ -925,12 +1146,48 @@ def convert_generation_model(args: argparse.Namespace, generation_type: Generati
|
|||
"Tried and failed to pad logits MatMul weights. " "Performance may be sub-optimal for this MatMul"
|
||||
)
|
||||
|
||||
gpt2_init_decoder_generated = False
|
||||
gpt2_init_decoder_onnx_path = None
|
||||
if (
|
||||
args.separate_gpt2_decoder_for_init_run
|
||||
and is_gpt2
|
||||
and (generation_type == GenerationType.BEAMSEARCH or generation_type == GenerationType.GREEDYSEARCH)
|
||||
):
|
||||
logger.info(f"Creating an initial run GPT2 decoder from {args.decoder_onnx}. ")
|
||||
|
||||
gpt2_init_decoder_onnx_filename = "gpt2_init_past_{}.onnx".format(
|
||||
"fp16" if args.precision == Precision.FLOAT16 else "fp32"
|
||||
)
|
||||
|
||||
gpt2_init_decoder_onnx_path = Path(Path(args.output).parent, gpt2_init_decoder_onnx_filename).as_posix()
|
||||
|
||||
gpt2_init_decoder_generated = generate_gpt2_init_decoder(
|
||||
args.decoder_onnx, gpt2_init_decoder_onnx_path, args.use_external_data_format
|
||||
)
|
||||
|
||||
if not gpt2_init_decoder_generated:
|
||||
logger.warning(
|
||||
"Tried and failed to generate the init decoder GPT2 model. "
|
||||
"Performance may be sub-optimal for the initial decoding run"
|
||||
)
|
||||
|
||||
# Update the graph input shapes for the non-initial decoder model to account
|
||||
# for the fact that the sequence length will always be 1
|
||||
if gpt2_init_decoder_generated and not update_input_shapes_for_gpt2_decoder_model(
|
||||
args.decoder_onnx, args.use_external_data_format
|
||||
):
|
||||
# Can't proceed further - better to raise an exception
|
||||
raise ValueError(f"Could not update the input shapes for the non-initial decoder subgraph.")
|
||||
|
||||
# If the user explicitly requests running shape inference or if we padded/mutated
|
||||
# weight(s) in the decoder, we want to run shape inference to capture the new
|
||||
# weight(s)/input shape(s) in the decoder, we want to run shape inference to capture the new
|
||||
# shapes
|
||||
if logits_matmul_weight_padded or args.run_shape_inference:
|
||||
if logits_matmul_weight_padded or args.run_shape_inference or gpt2_init_decoder_generated:
|
||||
logger.info(f"Run symbolic shape inference on {args.decoder_onnx}. The file will be overwritten.")
|
||||
shape_inference(args.decoder_onnx, args.use_external_data_format)
|
||||
if gpt2_init_decoder_generated:
|
||||
logger.info(f"Run symbolic shape inference on {gpt2_init_decoder_onnx_path}. The file will be overwritten.")
|
||||
shape_inference(gpt2_init_decoder_onnx_path, args.use_external_data_format)
|
||||
|
||||
if is_gpt2:
|
||||
config = GPT2Config.from_pretrained(args.model_name_or_path, cache_dir=args.cache_dir)
|
||||
|
|
@ -955,6 +1212,12 @@ def convert_generation_model(args: argparse.Namespace, generation_type: Generati
|
|||
|
||||
if args.model_type == "gpt2":
|
||||
verify_gpt2_subgraph(decoder_model.graph, args.precision)
|
||||
|
||||
# If we generated the init decoder model, verify that as well
|
||||
if gpt2_init_decoder_generated:
|
||||
gpt2_init_decoder_model = onnx.load_model(gpt2_init_decoder_onnx_path, load_external_data=True)
|
||||
gpt2_init_decoder_model.graph.name = f"{args.model_type} init decoder"
|
||||
verify_gpt2_subgraph(gpt2_init_decoder_model.graph, args.precision)
|
||||
else:
|
||||
verify_t5_decoder_subgraph(decoder_model.graph, args.precision)
|
||||
|
||||
|
|
@ -1054,7 +1317,7 @@ def convert_generation_model(args: argparse.Namespace, generation_type: Generati
|
|||
# Unique shared initializers from the decoder and decoder_init could reduce memory usage in inference.
|
||||
initializers = get_shared_initializers(encoder_model, decoder_model)
|
||||
logger.info(
|
||||
f"{len(initializers)} shared initializers ({[i.name for i in initializers]}) in subgraphs are moved to the main graph"
|
||||
f"{len(initializers)} shared initializers ({[i.name for i in initializers]}) in encoder and decoder subgraphs are moved to the main graph"
|
||||
)
|
||||
|
||||
# TODO(tianleiwu): investigate the following which causes error in inference
|
||||
|
|
@ -1076,9 +1339,23 @@ def convert_generation_model(args: argparse.Namespace, generation_type: Generati
|
|||
]
|
||||
)
|
||||
else:
|
||||
# Move initializer from subgraph to main graph could reduce memory usage in inference.
|
||||
initializers = move_initializers(decoder_model.graph)
|
||||
logger.info(f"{len(initializers)} initializers from the decoder are moved to the main graph")
|
||||
if gpt2_init_decoder_generated:
|
||||
gpt2_init_decoder_model = onnx.load_model(gpt2_init_decoder_onnx_path, load_external_data=True)
|
||||
|
||||
# Move shared initializers (shared between init decoder and decoder models) to the main
|
||||
# graph and remove them from these models
|
||||
if not args.disable_shared_initializers:
|
||||
# Unique shared initializers from the decoder and decoder_init could reduce memory usage in inference.
|
||||
initializers = get_shared_initializers(gpt2_init_decoder_model, decoder_model)
|
||||
logger.info(
|
||||
f"{len(initializers)} shared initializers ({[i.name for i in initializers]}) in decoder and init decoder subgraphs are moved to the main graph"
|
||||
)
|
||||
|
||||
node.attribute.append(onnx.helper.make_attribute("init_decoder", gpt2_init_decoder_model.graph))
|
||||
else:
|
||||
# Move initializer from subgraph to main graph could reduce memory usage in inference.
|
||||
initializers = move_initializers(decoder_model.graph)
|
||||
logger.info(f"{len(initializers)} initializers from the decoder are moved to the main graph")
|
||||
|
||||
node.attribute.append(onnx.helper.make_attribute("decoder", decoder_model.graph))
|
||||
|
||||
|
|
|
|||
|
|
@ -1041,7 +1041,8 @@ class OnnxModel:
|
|||
logger.warning("add_prefix_to_names does not process subgraphs.")
|
||||
|
||||
# Exclude the names of inputs and outputs of main graph (but not subgraphs)
|
||||
excluded = [i.name for i in self.model.graph.input] + [o.name for o in self.model.graph.output]
|
||||
# and empty names ("") as they have special meaning to denote missing optional inputs
|
||||
excluded = [i.name for i in self.model.graph.input] + [o.name for o in self.model.graph.output] + [""]
|
||||
|
||||
for initializer in self.model.graph.initializer:
|
||||
if initializer.name not in excluded:
|
||||
|
|
|
|||
Loading…
Reference in a new issue