【发布时间】:2021-10-25 06:32:31
【问题描述】:
下面是示例代码:
@deviceCountAtLeast(1)
if NO_DOUBLE:
@dtypes(torch.float)
else:
@dtypes(torch.float, torch.double)
def test_requires_grad_factory(self, devices, dtype):
fns = [torch.ones_like, torch.testing.randn_like]
x = torch.randn(2, 3, dtype=dtype, device=devices[0])
for fn in fns:
for requires_grad in [True, False]:
output = fn(x, dtype=dtype, device=devices[0], requires_grad=requires_grad)
self.assertEqual(requires_grad, output.requires_grad)
self.assertIs(dtype, output.dtype)
self.assertEqual(devices[0], str(x.device))
如您所见,我想根据NO_DOUBLE 值选择@dtypes() 装饰器的参数列表。
我目前的解决方法就像使用另一个函数来返回不同的装饰器:
def no_double(cond, dec1, dec2):
return dec1 if cond else dec2
@no_double(NO_DOUBLE, dtypes(torch.float), dtypes(torch.float, torch.double))
def test_requires_grad_factory(self, devices, dtype):
【问题讨论】: