MCPcopy Create free account
hub / github.com/MegEngine/MegEngine / assert_tensor_eq_with_iter

Function assert_tensor_eq_with_iter

dnn/test/common/checker.cpp:11–64  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

9namespace {
10template <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
66template <typename ctype>
67::testing::AssertionResult assert_tensor_eq_with_dtype(

Callers

nothing calls this directly

Calls 6

diffFunction · 0.85
good_floatFunction · 0.70
absFunction · 0.50
total_nr_elemsMethod · 0.45
to_stringMethod · 0.45
nameMethod · 0.45

Tested by

no test coverage detected