Skip to content

Commit b41f00c

Browse files
committed
Fix blas_nrm2 unittest
1 parent 43c49ad commit b41f00c

File tree

1 file changed

+5
-4
lines changed
  • source/source_base/module_container/ATen/kernels/test

1 file changed

+5
-4
lines changed

source/source_base/module_container/ATen/kernels/test/blas_test.cpp

Lines changed: 5 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -45,11 +45,12 @@ TYPED_TEST(BlasTest, Nrm2) {
4545
const int n = 3;
4646
const Tensor x = std::move(Tensor({static_cast<Type>(3.0), static_cast<Type>(4.0), static_cast<Type>(0.0)}).to_device<Device>());
4747

48-
Type result = {};
49-
nrm2Calculator(n, x.data<Type>(), 1, &result);
50-
const Type expected = static_cast<Type>(5.0);
48+
using Real = typename GetTypeReal<Type>::type;
49+
Real result = {};
50+
result = nrm2Calculator(n, x.data<Type>(), 1);
51+
const Real expected = static_cast<Real>(5.0);
5152

52-
EXPECT_NEAR(result, expected, static_cast<Type>(1e-6));
53+
EXPECT_NEAR(result, expected, static_cast<Real>(1e-6));
5354
}
5455

5556
TYPED_TEST(BlasTest, Dot) {

0 commit comments

Comments
 (0)