| 9 | namespace { |
| 10 | template <typename ctype, class Iter> |
| 11 | ::testing::AssertionResult assert_tensor_eq_with_iter( |
| 12 | const char* expr0, const char* expr1, Iter it0, Iter it1, |
| 13 | const TensorLayout& layout, float maxerr, float maxerr_avg, |
| 14 | float maxerr_avg_biased, bool allow_invalid) { |
| 15 | auto nr_elem = layout.total_nr_elems(); |
| 16 | double error_sum = 0; |
| 17 | double error_sum_biased = 0; |
| 18 | for (size_t i = 0; i < nr_elem; ++i) { |
| 19 | ctype iv0 = *it0, iv1 = *it1; |
| 20 | float err = diff(iv0, iv1); |
| 21 | error_sum += std::abs(err); |
| 22 | error_sum_biased += err; |
| 23 | if (!allow_invalid && |
| 24 | (!good_float(iv0) || !good_float(iv1) || std::abs(err) > maxerr)) { |
| 25 | Index index(layout, i); |
| 26 | return ::testing::AssertionFailure() |
| 27 | << "Unequal value\n" |
| 28 | << "Value of: " << expr1 << "\n" |
| 29 | << " Actual: " << (iv1 + 0) << "\n" |
| 30 | << "Expected: " << expr0 << "\n" |
| 31 | << "Which is: " << (iv0 + 0) << "\n" |
| 32 | << "At index: " << index.to_string() << "/" |
| 33 | << layout.TensorShape::to_string() << "\n" |
| 34 | << " DType: " << layout.dtype.name() << "\n" |
| 35 | << " error: " << std::abs(err) << "/" << maxerr; |
| 36 | } |
| 37 | |
| 38 | ++it0; |
| 39 | ++it1; |
| 40 | } |
| 41 | |
| 42 | float error_avg = error_sum / nr_elem; |
| 43 | if (error_avg > maxerr_avg) { |
| 44 | return ::testing::AssertionFailure() |
| 45 | << "Average error exceeds the upper limit\n" |
| 46 | << "Value of: " << expr1 << "\n" |
| 47 | << "Expected: " << expr0 << "\n" |
| 48 | << "Average error: " << error_avg << "/" << maxerr_avg << "\n" |
| 49 | << "Num of elements: " << nr_elem; |
| 50 | } |
| 51 | |
| 52 | float error_avg_biased = error_sum_biased / nr_elem; |
| 53 | if (std::abs(error_avg_biased) > maxerr_avg_biased) { |
| 54 | return ::testing::AssertionFailure() |
| 55 | << "Average biased error exceeds the upper limit\n" |
| 56 | << "Value of: " << expr1 << "\n" |
| 57 | << "Expected: " << expr0 << "\n" |
| 58 | << "Average biased error: " << error_avg_biased << "/" << maxerr_avg_biased |
| 59 | << "\n" |
| 60 | << "Num of elements: " << nr_elem; |
| 61 | } |
| 62 | |
| 63 | return ::testing::AssertionSuccess(); |
| 64 | } |
| 65 | |
| 66 | template <typename ctype> |
| 67 | ::testing::AssertionResult assert_tensor_eq_with_dtype( |
nothing calls this directly
no test coverage detected