Remove 'model_.' prefix from onnx model initializers in training (#3881)

* Remove 'model_.' prefix for onnx model initializers in training

* fix test case remove redundant device test

* rename

* Fix state_dict/load_state_dict with frozen_weight

* nit

* Add monkey patch for pt opset 10

* remove pt patch in CI

* nit: newline
This commit is contained in:
Bowen Bao 2020-05-20 10:06:31 -07:00 committed by GitHub
parent 08763e80e0
commit 0a5395bb78
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
6 changed files with 179 additions and 229 deletions

View file

@ -6,6 +6,7 @@ import unittest
import pytest
import sys
import copy
import numpy as np
from numpy.testing import assert_allclose, assert_array_equal
import onnx
@ -260,9 +261,9 @@ class MNISTWrapper():
model_desc = MNISTWrapper.mnist_model_description()
return model, model_desc
def get_trainer(self, model, model_desc, device, onnx_opset_ver=12):
def get_trainer(self, model, model_desc, device, onnx_opset_ver=12, frozen_weights=[]):
return ORTTrainer(model, MNISTWrapper.my_loss, model_desc, "SGDOptimizer", None, IODescription('Learning_Rate', [1, ],
torch.float32), device, _opset_version=onnx_opset_ver)
torch.float32), device, _opset_version=onnx_opset_ver, frozen_weights=frozen_weights)
class TestOrtTrainer(unittest.TestCase):
@ -386,7 +387,7 @@ class TestOrtTrainer(unittest.TestCase):
loss, _ = trainer.train_step(data, target, torch.tensor([learningRate]))
state_dict = trainer.state_dict()
assert state_dict.keys() == {'model_.fc1.bias', 'model_.fc1.weight', 'model_.fc2.bias', 'model_.fc2.weight'}
assert state_dict.keys() == {'fc1.bias', 'fc1.weight', 'fc2.bias', 'fc2.weight'}
def testMNISTSaveAsONNX(self):
torch.manual_seed(1)
@ -435,6 +436,93 @@ class TestOrtTrainer(unittest.TestCase):
loss, _ = trainer.train_step(data, target, torch.tensor([learningRate]))
def testMNISTInitializerNames(self):
torch.manual_seed(1)
device = torch.device("cuda")
mnist = MNISTWrapper()
train_loader, test_loader = mnist.get_loaders()
model, model_desc = mnist.get_model()
trainer = mnist.get_trainer(model, model_desc, device)
learningRate = 0.02
epoch = 0
data, target = next(iter(train_loader))
data, target = data.to(device), target.to(device)
data = data.reshape(data.shape[0], -1)
loss, _ = trainer.train_step(data, target, torch.tensor([learningRate]))
assert set([n.name for n in trainer.onnx_model_.graph.initializer]) \
== set([n for n, t in model.named_parameters()])
def testMNISTFrozenWeight(self):
torch.manual_seed(1)
device = torch.device("cuda")
mnist = MNISTWrapper()
train_loader, test_loader = mnist.get_loaders()
model, model_desc = mnist.get_model()
trainer = mnist.get_trainer(model, model_desc, device, frozen_weights=['fc1.weight'])
learningRate = 0.02
epoch = 0
data, target = next(iter(train_loader))
data, target = data.to(device), target.to(device)
data = data.reshape(data.shape[0], -1)
loss, _ = trainer.train_step(data, target, torch.tensor([learningRate]))
fc1_trainstep_1 = trainer.state_dict()['fc1.weight']
fc2_trainstep_1 = trainer.state_dict()['fc2.weight']
loss, _ = trainer.train_step(data, target, torch.tensor([learningRate]))
fc1_trainstep_2 = trainer.state_dict()['fc1.weight']
fc2_trainstep_2 = trainer.state_dict()['fc2.weight']
assert np.array_equal(fc1_trainstep_1, fc1_trainstep_2) and \
not np.array_equal(fc2_trainstep_1, fc2_trainstep_2)
def testMNISTFrozenWeightCheckpoint(self):
torch.manual_seed(1)
device = torch.device("cuda")
mnist = MNISTWrapper()
train_loader, test_loader = mnist.get_loaders()
model, model_desc = mnist.get_model()
trainer = mnist.get_trainer(model, model_desc, device, frozen_weights=['fc1.weight'])
learningRate = 0.02
epoch = 0
# do one train step
data, target = next(iter(train_loader))
data, target = data.to(device), target.to(device)
data = data.reshape(data.shape[0], -1)
loss, _ = trainer.train_step(data, target, torch.tensor([learningRate]))
# do one eval step
data, target = next(iter(train_loader))
data, target = data.to(device), target.to(device)
data = data.reshape(data.shape[0], -1)
loss, _ = trainer.eval_step(data, target)
# save checkpoint, load model and compare
state_dict = trainer.state_dict()
new_model, _ = mnist.get_model()
trainer = mnist.get_trainer(new_model, model_desc, device, frozen_weights=['fc1.weight'])
trainer.load_state_dict(state_dict)
ckpt_loss, _ = trainer.eval_step(data, target)
assert loss == ckpt_loss
def testBertTrainingBasic(self):
expected_losses = [
11.02906322479248, 11.094074249267578, 11.00899887084961, 11.06129264831543,

Binary file not shown.

View file

@ -11,6 +11,7 @@ import torch.onnx
import onnxruntime as ort
from distutils.version import LooseVersion
from .checkpointing_utils import list_checkpoint_files, get_checkpoint_name, CombineZeroCheckpoint
import onnxruntime.capi.pt_patch
DEFAULT_OPSET_VERSION = 10
@ -326,11 +327,29 @@ def convert_model_loss_fn_to_onnx(model, loss_fn, model_desc, device, inputs, op
do_constant_folding=False,
**other_export_options)
model = onnx.load_model_from_string(f.getvalue())
onnx_model = onnx.load_model_from_string(f.getvalue())
model = FuseSofmaxNLLToSoftmaxCE(model)
# Remove 'model_.' prefix introduced by model wrapper for initializers.
replace_name_dict = {}
for n in onnx_model.graph.initializer:
if n.name.startswith('model_.'):
replace_name_dict[n.name] = n.name[len('model_.'):]
n.name = replace_name_dict[n.name]
for n in onnx_model.graph.node:
for i, name in enumerate(n.input):
if name in replace_name_dict:
n.input[i] = replace_name_dict[name]
return model
# onnx model initializer may contain non-trainable registered buffers that are not part
# of pytorch model named parameteres.
assert set([n for n, t in model.model_.named_parameters()]).issubset(
set([n.name for n in onnx_model.graph.initializer])), \
"Initializer names do not match between PyTorch model and ONNX model, " \
"please report a bug to ONNX Runtime."
onnx_model = FuseSofmaxNLLToSoftmaxCE(onnx_model)
return onnx_model
def create_ort_training_session_with_optimizer(model, device, training_optimizer_name, lr_params_feed_name,
map_optimizer_attributes, world_rank=-1, world_size=1,
@ -361,6 +380,11 @@ def create_ort_training_session_with_optimizer(model, device, training_optimizer
torch_params = {}
optimizer_attributes_map = {}
optimizer_int_attributes_map = {}
unused_frozen_weights = [n for n in frozen_weights if n not in [i.name for i in model.graph.initializer]]
if unused_frozen_weights:
raise RuntimeError("{} in frozen_weights not found in model weights.".format(unused_frozen_weights))
weights_to_train = set()
for initializer in model.graph.initializer:
if initializer.name in frozen_weights:
@ -639,15 +663,36 @@ class ORTTrainer():
def eval(self):
self.is_train = False
def _update_onnx_model_initializers(self, state_tensors):
# replace the initializers with new value
new_weights = []
replace_indices = []
for i, w in enumerate(self.onnx_model_.graph.initializer):
if w.name in state_tensors:
new_weights.append(numpy_helper.from_array(state_tensors[w.name], w.name))
replace_indices.append(i)
replace_indices.sort(reverse=True)
for w_i in replace_indices:
del self.onnx_model_.graph.initializer[w_i]
self.onnx_model_.graph.initializer.extend(new_weights)
def state_dict(self):
if not self.session:
warnings.warn("ONNXRuntime training session is not initialized yet. "
"Please run train_step or eval_step at least once before calling state_dict().")
return {}
# extract trained weights
session_state = self.session.get_state()
torch_state = {}
for name in session_state:
torch_state[name] = torch.from_numpy(session_state[name])
# extract untrained weights and buffer
for n in self.onnx_model_.graph.initializer:
if n.name not in torch_state:
torch_state[n.name] = torch.from_numpy(numpy_helper.to_array(n))
return torch_state
def load_state_dict(self, state_dict, strict=False):
@ -659,10 +704,21 @@ class ORTTrainer():
self.strict_ = strict
return
session_state = {}
# update onnx model from loaded state dict
cur_initializers_names = [n.name for n in self.onnx_model_.graph.initializer]
new_initializers = {}
for name in state_dict:
session_state[name] = state_dict[name].numpy()
self.session.load_state(session_state, strict)
if name in cur_initializers_names:
new_initializers[name] = state_dict[name].numpy()
elif strict:
raise RuntimeError("Checkpoint tensor: {} is not present in the model.".format(name))
self._update_onnx_model_initializers(new_initializers)
# create new session based on updated onnx model
self.state_dict_ = None
self._init_session()
def save_as_onnx(self, path):
if not self.session:
@ -670,19 +726,7 @@ class ORTTrainer():
"Please run train_step or eval_step at least once before calling save_as_onnx().")
return
state_tensors = self.session.get_state()
# replace the initializers with new value
new_weights = []
replace_indices = []
i = 0
for w in self.onnx_model_.graph.initializer:
if w.name in state_tensors:
new_weights.append(numpy_helper.from_array(state_tensors[w.name], w.name))
replace_indices.append(i)
i += 1
replace_indices.sort(reverse=True)
for w_i in replace_indices:
del self.onnx_model_.graph.initializer[w_i]
self.onnx_model_.graph.initializer.extend(new_weights)
self._update_onnx_model_initializers(state_tensors)
with open(path, "wb") as f:
f.write(self.onnx_model_.SerializeToString())

View file

@ -0,0 +1,21 @@
import torch
from torch.onnx import symbolic_opset10
from torch.onnx.symbolic_helper import parse_args
@parse_args('v', 'v', 'v', 'v', 'i', 'none')
def nll_loss(g, self, target, weight=None, reduction='mean', ignore_index=-100):
if not weight and not ignore_index:
return g.op("nll_loss", self, target)
elif ignore_index:
ignore_index_ = g.op("Constant", value_t=torch.tensor(ignore_index, dtype=torch.int64))
eq_ = g.op("Equal", target, ignore_index_)
not_eq_ = g.op("Not", eq_)
weight_ = g.op("Cast", not_eq_, to_i=1) # FLOAT = 1; // float
not_eq_int64_ = g.op("Cast", not_eq_, to_i=7) #INT64 = 7; // int64_t
target_ = g.op("Mul", target, not_eq_int64_)
# if weight:
# weight_ = g.op("Mul", weight_, weight)
return g.op("nll_loss", self, target_, weight_)
symbolic_opset10.nll_loss = nll_loss

View file

@ -44,7 +44,7 @@ function GetFile {
if command -v aria2c > /dev/null; then
aria2c -q -d $(dirname $path) -o $(basename $path) "$uri"
else
curl "$uri" -sSL --retry $download_retries --retry-delay $retry_wait_time_seconds --create-dirs -o "$path" --fail
curl "$uri" -sSL --retry $download_retries --retry-delay $retry_wait_time_seconds --create-dirs -o "$path" --fail
fi
return $?
@ -76,15 +76,15 @@ fi
if [[ $SYS_LONG_BIT = "64" && "$GLIBC_VERSION" -gt "9" ]]; then
echo "Installing azcopy"
mkdir -p /tmp/azcopy
GetFile https://aka.ms/downloadazcopy-v10-linux /tmp/azcopy/azcopy.tar.gz
GetFile https://aka.ms/downloadazcopy-v10-linux /tmp/azcopy/azcopy.tar.gz
tar --strip 1 -xf /tmp/azcopy/azcopy.tar.gz -C /tmp/azcopy
cp /tmp/azcopy/azcopy /usr/bin
echo "Installing cmake"
GetFile https://github.com/Kitware/CMake/releases/download/v3.13.5/cmake-3.13.5-Linux-x86_64.tar.gz /tmp/src/cmake-3.13.5-Linux-x86_64.tar.gz
GetFile https://github.com/Kitware/CMake/releases/download/v3.13.5/cmake-3.13.5-Linux-x86_64.tar.gz /tmp/src/cmake-3.13.5-Linux-x86_64.tar.gz
tar -zxf /tmp/src/cmake-3.13.5-Linux-x86_64.tar.gz --strip=1 -C /usr
else
echo "Installing cmake"
GetFile https://github.com/Kitware/CMake/releases/download/v3.13.5/cmake-3.13.5.tar.gz /tmp/src/cmake-3.13.5.tar.gz
GetFile https://github.com/Kitware/CMake/releases/download/v3.13.5/cmake-3.13.5.tar.gz /tmp/src/cmake-3.13.5.tar.gz
tar -xf /tmp/src/cmake-3.13.5.tar.gz -C /tmp/src
pushd .
cd /tmp/src/cmake-3.13.5
@ -112,10 +112,6 @@ elif [ $DEVICE_TYPE = "gpu" ]; then
${PYTHON_EXE} -m pip install sympy==1.1.1
if [[ $BUILD_EXTR_PAR = *--enable_training* ]]; then
${PYTHON_EXE} -m pip install --upgrade --pre torch torchvision -f https://download.pytorch.org/whl/nightly/cu101/torch_nightly.html
# patch pytorch onnx export opset version 10 to export nll_loss
PATH_TO_SYMBOLIC10=$(${PYTHON_EXE} -c 'import torch; import os; print(os.path.join(os.path.dirname(torch.__file__), "onnx/"))')
cp "${SCRIPT_DIR}/pyt_patch/symbolic_opset10.py" "${PATH_TO_SYMBOLIC10}"
fi
if [[ $BUILD_EXTR_PAR = *--enable_training_python_frontend_e2e_tests* ]]; then
${PYTHON_EXE} -m pip install transformers

View file

@ -1,199 +0,0 @@
from __future__ import absolute_import, division, print_function, unicode_literals
import torch
from torch.nn.modules.utils import _single, _pair, _triple
import torch.onnx
# This import monkey-patches graph manipulation methods on Graph, used for the
# ONNX symbolics
import torch.onnx.utils
import torch.onnx.symbolic_helper as sym_help
from torch.onnx.symbolic_helper import parse_args, _unimplemented
import torch.onnx.symbolic_opset9
# EDITING THIS FILE? READ THIS FIRST!
# see Note [Edit Symbolic Files] in symbolic_helper.py
# This file exports ONNX ops for opset 10
# Opset 10 is supported by ONNX release 1.5.0
# release on 04/24/19
@parse_args('v', 'i', 'i', 'none')
def sort(g, self, dim, decending, out=None):
return sym_help._sort_helper(g, self, dim, decending=decending, out=out)
@parse_args('v', 'v', 'i', 'i', 'i', 'none')
def topk(g, self, k, dim, largest, sorted, out=None):
return sym_help._topk_helper(g, self, k, dim, largest=largest, sorted=sorted, out=out)
def _max_pool(name, tuple_fn, ndims, return_indices):
@parse_args('v', 'is', 'is', 'is', 'is', 'i')
def symbolic_fn(g, input, kernel_size, stride, padding, dilation, ceil_mode):
if not stride:
stride = kernel_size
kwargs = {
'kernel_shape_i': tuple_fn(kernel_size),
'pads_i': tuple_fn(padding) * 2,
'strides_i': tuple_fn(stride),
'ceil_mode_i': ceil_mode,
}
if set(tuple_fn(dilation)) != {1}:
kwargs['dilations_i'] = tuple_fn(dilation)
# easy but hacky way to get flattened indices values
# to be used to convert the indices values to non-flattened.
# In ONNX the indices are computed as a flatten 1-D tensor,
# so the values in indices are in [0, N x C x D1 x ... x Dn).
# To convert the indices to the same format used by Pytorch,
# we first execute a maxpool with a kernel and stride of 1 on the same input.
# This will result in a tensor of indices in which each index will have it's own value.
# Using this tensor as a reference, we extract the first index of each axis and subtract
# it from each index of this axis in the indices to convert.
# This step will result in a tensor were each dimension has values of indices within
# the dimension it is in.
# For more information :
# https://github.com/pytorch/pytorch/pull/16455#issuecomment-460776407
if return_indices:
r, indices = g.op("MaxPool", input, outputs=2, **kwargs)
_, flattened_indices = g.op("MaxPool", input, outputs=2,
kernel_shape_i=[1 for _ in range(ndims)],
strides_i=[1 for _ in range(ndims)])
# convert indices to have non-flattened indices values
from torch.onnx.symbolic_opset9 import sub
s = sym_help._slice_helper(g, flattened_indices, axes=[2 + i for i in range(ndims)],
starts=tuple_fn(0), ends=tuple_fn(1))
indices = sub(g, indices, s)
return r, indices
else:
r = g.op("MaxPool", input, outputs=1, **kwargs)
return r
return symbolic_fn
max_pool1d = _max_pool("max_pool1d", _single, 1, return_indices=False)
max_pool2d = _max_pool("max_pool2d", _pair, 2, return_indices=False)
max_pool3d = _max_pool("max_pool3d", _triple, 3, return_indices=False)
max_pool1d_with_indices = _max_pool("max_pool1d_with_indices", _single, 1, return_indices=True)
max_pool2d_with_indices = _max_pool("max_pool2d_with_indices", _pair, 2, return_indices=True)
max_pool3d_with_indices = _max_pool("max_pool3d_with_indices", _triple, 3, return_indices=True)
def _avg_pool(name, tuple_fn):
@parse_args('v', 'is', 'is', 'is', 'i', 'i', 'none')
def symbolic_fn(g, input, kernel_size, stride, padding, ceil_mode, count_include_pad, divisor_override=None):
if not stride:
stride = kernel_size
padding = sym_help._avgpool_helper(tuple_fn, padding, kernel_size, stride, divisor_override, name)
if count_include_pad:
input = g.op("Pad", input,
pads_i=((0,) * 2 + padding) * 2,
mode_s='constant',
value_f=0.)
padding = (0,) * len(padding)
output = g.op("AveragePool", input,
kernel_shape_i=tuple_fn(kernel_size),
strides_i=tuple_fn(stride),
pads_i=padding * 2,
ceil_mode_i=ceil_mode)
return output
return symbolic_fn
avg_pool1d = _avg_pool('avg_pool1d', _single)
avg_pool2d = _avg_pool('avg_pool2d', _pair)
avg_pool3d = _avg_pool('avg_pool3d', _triple)
def _interpolate(name, dim, interpolate_mode):
def symbolic_fn(g, input, output_size, *args):
scales, align_corners = sym_help._get_interpolate_attributes(g, interpolate_mode, args)
sym_help._interpolate_warning(interpolate_mode)
align_corners = sym_help._maybe_get_scalar(align_corners)
if align_corners:
return _unimplemented(name, "align_corners == True")
if scales is None:
scales = sym_help._interpolate_size_to_scales(g, input, output_size, dim)
return g.op("Resize", input, scales, mode_s=interpolate_mode)
return symbolic_fn
upsample_nearest1d = _interpolate('upsample_nearest1d', 3, "nearest")
upsample_nearest2d = _interpolate('upsample_nearest2d', 4, "nearest")
upsample_nearest3d = _interpolate('upsample_nearest3d', 5, "nearest")
upsample_linear1d = _interpolate('upsample_linear1d', 3, "linear")
upsample_bilinear2d = _interpolate('upsample_bilinear2d', 4, "linear")
upsample_trilinear3d = _interpolate('upsample_trilinear3d', 5, "linear")
def __interpolate(g, input, size, scale_factor, mode, align_corners, recompute_scale_factor):
scales, mode = sym_help._interpolate_get_scales_and_mode(g, input, size, scale_factor,
mode, align_corners)
return g.op("Resize", input, scales, mode_s=mode)
def _slice(g, input, axes, starts, ends, steps=None, dynamic_slice=False):
if dynamic_slice:
starts = g.op("Unsqueeze", starts, axes_i=[0])
ends = g.op("Unsqueeze", ends, axes_i=[0])
axes = g.op("Unsqueeze", axes, axes_i=[0])
else:
assert len(starts) == len(ends)
assert len(starts) == len(axes)
assert steps is None or len(starts) == len(steps)
if len(starts) == 1 and starts[0] == 0 and ends[0] == 9223372036854775807 \
and (steps is None or (len(steps) == 1 and steps[0] == 1)):
return input
axes = g.op("Constant", value_t=torch.tensor(axes))
starts = g.op("Constant", value_t=torch.tensor(starts))
ends = g.op("Constant", value_t=torch.tensor(ends))
if steps is None:
return g.op("Slice", input, starts, ends, axes)
steps = g.op("Constant", value_t=torch.tensor(steps))
return g.op("Slice", input, starts, ends, axes, steps)
@parse_args('v', 'v', 'v', 'v', 'i')
def slice(g, self, dim, start, end, step):
if (start.node().kind() != 'onnx::Constant' or
end.node().kind() != 'onnx::Constant' or dim.node().kind() != 'onnx::Constant'):
dynamic_slice = True
else:
start = [sym_help._parse_arg(start, 'i')]
end = [sym_help._parse_arg(end, 'i')]
dim = [sym_help._parse_arg(dim, 'i')]
dynamic_slice = False
return sym_help._slice_helper(g, self, axes=dim, starts=start, ends=end, steps=[step], dynamic_slice=dynamic_slice)
@parse_args('v', 'is')
def flip(g, input, dims):
return sym_help._slice_helper(g, input, axes=dims,
starts=[-1] * len(dims),
ends=[-9223372036854775807] * len(dims),
steps=[-1] * len(dims))
def fmod(g, input, other):
return g.op("Mod", input, other, fmod_i=1)
# put nll in version 10 because ORT does not support some ops (like Equal) beyound opset 10.
@parse_args('v', 'v', 'v', 'v', 'i', 'none')
def nll_loss(g, self, target, weight=None, reduction='mean', ignore_index=-100):
if not weight and not ignore_index:
return g.op("nll_loss", self, target)
elif ignore_index:
ignore_index_ = g.op("Constant", value_t=torch.tensor(ignore_index, dtype=torch.int64))
eq_ = g.op("Equal", target, ignore_index_)
not_eq_ = g.op("Not", eq_)
weight_ = g.op("Cast", not_eq_, to_i=1) # FLOAT = 1; // float
not_eq_int64_ = g.op("Cast", not_eq_, to_i=7) # INT64 = 7; // int64_t
target_ = g.op("Mul", target, not_eq_int64_)
# if weight:
# weight_ = g.op("Mul", weight_, weight)
return g.op("nll_loss", self, target_, weight_)