| 66 | } |
| 67 | |
| 68 | bool MatchAndExplainTensor(const Tensor& tensor, const Tensor& expected_tensor, |
| 69 | ::testing::MatchResultListener* listener) { |
| 70 | if (tensor.dtype() != expected_tensor.dtype()) { |
| 71 | if (listener->IsInterested()) { |
| 72 | *listener << "\nexpected tensor of type " |
| 73 | << DataType_Name(expected_tensor.dtype()) |
| 74 | << " but found one of type " << DataType_Name(tensor.dtype()); |
| 75 | return false; |
| 76 | } |
| 77 | } |
| 78 | |
| 79 | switch (tensor.dtype()) { |
| 80 | case DT_HALF: |
| 81 | return CompareTensor<Eigen::half>(tensor, expected_tensor, listener); |
| 82 | case DT_FLOAT: |
| 83 | return CompareTensor<float>(tensor, expected_tensor, listener); |
| 84 | case DT_DOUBLE: |
| 85 | return CompareTensor<double>(tensor, expected_tensor, listener); |
| 86 | case DT_INT8: |
| 87 | return CompareTensor<int8>(tensor, expected_tensor, listener); |
| 88 | case DT_INT16: |
| 89 | return CompareTensor<int16>(tensor, expected_tensor, listener); |
| 90 | case DT_INT32: |
| 91 | return CompareTensor<int32>(tensor, expected_tensor, listener); |
| 92 | case DT_INT64: |
| 93 | return CompareTensor<int64>(tensor, expected_tensor, listener); |
| 94 | case DT_UINT8: |
| 95 | return CompareTensor<uint8>(tensor, expected_tensor, listener); |
| 96 | case DT_UINT16: |
| 97 | return CompareTensor<uint16>(tensor, expected_tensor, listener); |
| 98 | case DT_UINT32: |
| 99 | return CompareTensor<uint32>(tensor, expected_tensor, listener); |
| 100 | case DT_UINT64: |
| 101 | return CompareTensor<uint64>(tensor, expected_tensor, listener); |
| 102 | default: |
| 103 | LOG(FATAL) << "Unsupported dtype " // Crash ok: testonly. |
| 104 | << DataType_Name(tensor.dtype()); |
| 105 | } |
| 106 | } |
| 107 | |
| 108 | struct NodeMatcher : public ::testing::MatcherInterface<const Node*> { |
| 109 | bool MatchAndExplain( |