diff --git a/onnxruntime/test/python/onnxruntime_test_ort_trainer.py b/onnxruntime/test/python/onnxruntime_test_ort_trainer.py index 98c9b6b899..3123ad2e89 100644 --- a/onnxruntime/test/python/onnxruntime_test_ort_trainer.py +++ b/onnxruntime/test/python/onnxruntime_test_ort_trainer.py @@ -178,11 +178,13 @@ class MNISTWrapper(): self.fc1 = nn.Linear(input_size, hidden_size) self.relu = nn.ReLU() self.fc2 = nn.Linear(hidden_size, num_classes) + self.register_buffer("bias_buffer", torch.tensor(1e-6)) def forward(self, x): out = self.fc1(x) out = self.relu(out) out = self.fc2(out) + out = torch.add(out, self.bias_buffer.to(out.dtype)) return out def my_loss(x, target): @@ -387,7 +389,7 @@ class TestOrtTrainer(unittest.TestCase): loss, _ = trainer.train_step(data, target, torch.tensor([learningRate])) state_dict = trainer.state_dict() - assert state_dict.keys() == {'fc1.bias', 'fc1.weight', 'fc2.bias', 'fc2.weight'} + assert state_dict.keys() == {'fc1.bias', 'fc1.weight', 'fc2.bias', 'fc2.weight', 'bias_buffer'} def testMNISTSaveAsONNX(self): torch.manual_seed(1) @@ -454,7 +456,7 @@ class TestOrtTrainer(unittest.TestCase): loss, _ = trainer.train_step(data, target, torch.tensor([learningRate])) - assert set([n.name for n in trainer.onnx_model_.graph.initializer]) \ + 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 testMNISTFrozenWeight(self): @@ -486,6 +488,35 @@ class TestOrtTrainer(unittest.TestCase): assert np.array_equal(fc1_trainstep_1, fc1_trainstep_2) and \ not np.array_equal(fc2_trainstep_1, fc2_trainstep_2) + def testMNISTTorchBuffer(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])) + + fc1_trainstep_1 = trainer.state_dict()['fc1.weight'] + bias_buffer_trainstep_1 = trainer.state_dict()['bias_buffer'] + + loss, _ = trainer.train_step(data, target, torch.tensor([learningRate])) + + fc1_trainstep_2 = trainer.state_dict()['fc1.weight'] + bias_buffer_trainstep_2 = trainer.state_dict()['bias_buffer'] + assert not np.array_equal(fc1_trainstep_1, fc1_trainstep_2) and \ + np.array_equal(bias_buffer_trainstep_1, bias_buffer_trainstep_2) + def testMNISTFrozenWeightCheckpoint(self): torch.manual_seed(1) device = torch.device("cuda") diff --git a/orttraining/orttraining/python/ort_trainer.py b/orttraining/orttraining/python/ort_trainer.py index 4d2d614b47..0115ab33fe 100644 --- a/orttraining/orttraining/python/ort_trainer.py +++ b/orttraining/orttraining/python/ort_trainer.py @@ -685,6 +685,9 @@ class ORTTrainer(): if self.torch_model_ is not None: # NOTE: pt model is moved to cpu to conserve gpu memory. 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.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_)