Co-authored-by: Brayden Zhong <brayden@radixark.ai>
This commit is contained in:
co-authored by
Brayden Zhong
parent
4d06d4c97f
commit
9495737d82
@@ -109,15 +109,15 @@ class TestTopkPaddedRegion(CustomTestCase):
|
|||||||
|
|
||||||
def test_invalid_pad_count_tensor_raises(self):
|
def test_invalid_pad_count_tensor_raises(self):
|
||||||
x = torch.rand((8, 8), device=self.DEVICE, dtype=torch.float32)
|
x = torch.rand((8, 8), device=self.DEVICE, dtype=torch.float32)
|
||||||
with self.assertRaises(AssertionError):
|
with self.assertRaises(TypeError):
|
||||||
_fill_padded_rows(x, 4, 0.0) # python int, not a tensor
|
_fill_padded_rows(x, 4, 0.0) # python int, not a tensor
|
||||||
with self.assertRaises(AssertionError):
|
with self.assertRaises(ValueError):
|
||||||
_fill_padded_rows(
|
_fill_padded_rows(
|
||||||
x,
|
x,
|
||||||
torch.tensor([1, 2], device=self.DEVICE, dtype=torch.int32),
|
torch.tensor([1, 2], device=self.DEVICE, dtype=torch.int32),
|
||||||
0.0,
|
0.0,
|
||||||
)
|
)
|
||||||
with self.assertRaises(AssertionError):
|
with self.assertRaises(TypeError):
|
||||||
_fill_padded_rows(
|
_fill_padded_rows(
|
||||||
x,
|
x,
|
||||||
torch.tensor(4.0, device=self.DEVICE, dtype=torch.float32),
|
torch.tensor(4.0, device=self.DEVICE, dtype=torch.float32),
|
||||||
|
|||||||
Reference in New Issue
Block a user