Merge nucleic/sleek-ember-seal-uady into dev
This commit is contained in:
@@ -12,6 +12,32 @@ import train
|
||||
|
||||
|
||||
class MetricsTests(unittest.TestCase):
|
||||
def test_qat_replacements_keep_checkpoint_keys_and_gradients(self):
|
||||
model = torch.nn.Sequential(
|
||||
torch.nn.Embedding(16, 8),
|
||||
torch.nn.Flatten(),
|
||||
torch.nn.Linear(16, 3),
|
||||
)
|
||||
keys = set(model.state_dict())
|
||||
counts = train.enable_quantization_aware_training(torch, model)
|
||||
self.assertEqual({"linear": 1, "embedding": 1}, counts)
|
||||
self.assertEqual(keys, set(model.state_dict()))
|
||||
output = model(torch.tensor([[1, 2]], dtype=torch.long))
|
||||
output.sum().backward()
|
||||
self.assertIsNotNone(model[0].weight.grad)
|
||||
self.assertIsNotNone(model[2].weight.grad)
|
||||
|
||||
def test_qat_affine_ranges_include_zero(self):
|
||||
model = torch.nn.Sequential(torch.nn.Linear(2, 2, bias=False))
|
||||
with torch.no_grad():
|
||||
model[0].weight.copy_(torch.eye(2))
|
||||
train.enable_quantization_aware_training(torch, model)
|
||||
output = model(torch.tensor([[1.0, 2.0]]))
|
||||
self.assertTrue(
|
||||
torch.allclose(output, torch.tensor([[1.0, 2.0]]), atol=0.02),
|
||||
output,
|
||||
)
|
||||
|
||||
def test_boundary_training_weight_is_opt_in(self):
|
||||
self.assertEqual(
|
||||
2.0,
|
||||
|
||||
Reference in New Issue
Block a user