| 18 | #include "paddle/phi/common/bfloat16.h" |
| 19 | |
| 20 | TEST(DataTransform, DataLayoutFunction) { |
| 21 | auto place = phi::CPUPlace(); |
| 22 | phi::DenseTensor in = phi::DenseTensor(); |
| 23 | phi::DenseTensor out = phi::DenseTensor(); |
| 24 | in.mutable_data<double>(common::make_ddim({2, 3, 1, 2}), place); |
| 25 | in.set_layout(phi::DataLayout::NHWC); |
| 26 | |
| 27 | auto kernel_nhwc = |
| 28 | phi::KernelKey(place, phi::DataLayout::NHWC, phi::DataType::FLOAT32); |
| 29 | auto kernel_nchw = |
| 30 | phi::KernelKey(place, phi::DataLayout::NCHW, phi::DataType::FLOAT32); |
| 31 | |
| 32 | paddle::framework::TransDataLayout(kernel_nhwc, kernel_nchw, in, &out, place); |
| 33 | |
| 34 | EXPECT_TRUE(out.layout() == phi::DataLayout::NCHW); |
| 35 | EXPECT_TRUE(out.dims() == common::make_ddim({2, 2, 3, 1})); |
| 36 | |
| 37 | paddle::framework::TransDataLayout(kernel_nchw, kernel_nhwc, in, &out, place); |
| 38 | |
| 39 | EXPECT_TRUE(in.layout() == phi::DataLayout::NHWC); |
| 40 | EXPECT_TRUE(in.dims() == common::make_ddim({2, 3, 1, 2})); |
| 41 | } |
| 42 | |
| 43 | #ifdef PADDLE_WITH_DNNL |
| 44 | TEST(DataTransformBf16, GetDataFromTensorDNNL) { |
nothing calls this directly
no test coverage detected