Bug fix for parameter names of models not using wrapper (#4061)

* bug fix for models not using wrapper

* add test case for no wrapper case

* update test case to use internal learning rate

* fix bug with frozen weight update
This commit is contained in:
Bowen Bao 2020-06-05 12:03:38 -07:00 committed by GitHub
parent 9790e19424
commit 1e5307d458
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
2 changed files with 59 additions and 6 deletions

View file

@ -78,7 +78,7 @@ def create_ort_trainer(gradient_accumulation_steps,
model = ORTTrainer(onnx_model, None, simple_model_desc, "LambOptimizer",
map_optimizer_attributes,
learning_rate_description,
device, postprocess_model=None,
device,
gradient_accumulation_steps=gradient_accumulation_steps,
world_rank=0, world_size=1,
loss_scaler=loss_scaler,
@ -187,6 +187,19 @@ class MNISTWrapper():
out = torch.add(out, self.bias_buffer.to(out.dtype))
return out
class NeuralNetWithLoss(nn.Module):
def __init__(self, input_size, hidden_size, num_classes):
super(MNISTWrapper.NeuralNetWithLoss, self).__init__()
self.fc1 = nn.Linear(input_size, hidden_size)
self.relu = nn.ReLU()
self.fc2 = nn.Linear(hidden_size, num_classes)
def forward(self, x, target):
out = self.fc1(x)
out = self.relu(out)
out = self.fc2(out)
return F.nll_loss(F.log_softmax(out, dim=1), target), out
def my_loss(x, target):
return F.nll_loss(F.log_softmax(x, dim=1), target)
@ -263,9 +276,22 @@ class MNISTWrapper():
model_desc = MNISTWrapper.mnist_model_description()
return model, model_desc
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, frozen_weights=frozen_weights)
def get_model_with_internal_loss(self):
input_size = 784
hidden_size = 500
num_classes = 10
# warning: changes the pytorch random generator state
model = MNISTWrapper.NeuralNetWithLoss(input_size, hidden_size, num_classes)
model_desc = MNISTWrapper.mnist_model_description()
return model, model_desc
def get_trainer(self, model, model_desc, device, onnx_opset_ver=12, frozen_weights=[],
internal_loss_fn=False, get_lr_this_step=None):
loss_fn = MNISTWrapper.my_loss if not internal_loss_fn else None
return ORTTrainer(model, loss_fn, model_desc, "SGDOptimizer", None, IODescription('Learning_Rate', [1, ],
torch.float32), device, _opset_version=onnx_opset_ver, frozen_weights=frozen_weights,
get_lr_this_step=get_lr_this_step)
class TestOrtTrainer(unittest.TestCase):
@ -459,6 +485,32 @@ class TestOrtTrainer(unittest.TestCase):
assert (set([n.name for n in trainer.onnx_model_.graph.initializer])-set(['bias_buffer'])) \
== set([n for n, t in model.named_parameters()])
def testMNISTInitializerNamesWithInternalLoss(self):
torch.manual_seed(1)
device = torch.device("cuda")
mnist = MNISTWrapper()
train_loader, test_loader = mnist.get_loaders()
model, model_desc = mnist.get_model_with_internal_loss()
def get_lr_this_step(global_step):
learningRate = 0.02
return torch.tensor([learningRate])
trainer = mnist.get_trainer(model, model_desc, device, internal_loss_fn=True,
get_lr_this_step=get_lr_this_step)
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)
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")

View file

@ -342,7 +342,8 @@ def convert_model_loss_fn_to_onnx(model, loss_fn, model_desc, device, inputs, op
# 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(
named_parameters = model.model_.named_parameters() if hasattr(model, 'model_') else model.named_parameters()
assert set([n for n, t in 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."
@ -687,7 +688,7 @@ class ORTTrainer():
self.torch_model_.cpu()
# torch buffers created using 'register_buffer' are not meant to be trainable.
torch_buffers = list(dict(self.torch_model_.named_buffers()).keys())
self.frozen_weights_.extend(torch_buffers)
self.frozen_weights_ = self.frozen_weights_ + torch_buffers
self.onnx_model_ = convert_model_loss_fn_to_onnx(
self.torch_model_, self.loss_fn_, self.model_desc_, torch.device('cpu'), inputs, opset_version=self.opset_version_)