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

Function assert_tensor_eq_with_iter

compiler/test/kernel/common/src/checker.cpp:141–193  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

139
140template <typename ctype, class Iter>
141::testing::AssertionResult assert_tensor_eq_with_iter(
142 const char* expr0, const char* expr1, Iter it0, Iter it1,
143 const TensorLayout& layout, float maxerr, float maxerr_avg,
144 float maxerr_avg_biased) {
145 auto nr_elem = layout.total_nr_elems();
146 double error_sum = 0;
147 double error_sum_biased = 0;
148 for (size_t i = 0; i < nr_elem; ++i) {
149 ctype iv0 = *it0, iv1 = *it1;
150 float err = diff(iv0, iv1);
151 error_sum += std::abs(err);
152 error_sum_biased += err;
153 if (!good_float(iv0) || !good_float(iv1) || std::abs(err) > maxerr) {
154 Index index(layout, i);
155 return ::testing::AssertionFailure()
156 << "Unequal value\n"
157 << "Value of: " << expr1 << "\n"
158 << " Actual: " << (iv1 + 0) << "\n"
159 << "Expected: " << expr0 << "\n"
160 << "Which is: " << (iv0 + 0) << "\n"
161 << "At index: " << index.to_string() << "/"
162 << layout.TensorShape::to_string() << "\n"
163 << " DType: " << layout.dtype.name() << "\n"
164 << " error: " << std::abs(err) << "/" << maxerr;
165 }
166
167 ++it0;
168 ++it1;
169 }
170
171 float error_avg = error_sum / nr_elem;
172 if (error_avg > maxerr_avg) {
173 return ::testing::AssertionFailure()
174 << "Average error exceeds the upper limit\n"
175 << "Value of: " << expr1 << "\n"
176 << "Expected: " << expr0 << "\n"
177 << "Average error: " << error_avg << "/" << maxerr_avg << "\n"
178 << "Num of elements: " << nr_elem;
179 }
180
181 float error_avg_biased = error_sum_biased / nr_elem;
182 if (std::abs(error_avg_biased) > maxerr_avg_biased) {
183 return ::testing::AssertionFailure()
184 << "Average biased error exceeds the upper limit\n"
185 << "Value of: " << expr1 << "\n"
186 << "Expected: " << expr0 << "\n"
187 << "Average biased error: " << error_avg_biased << "/" << maxerr_avg_biased
188 << "\n"
189 << "Num of elements: " << nr_elem;
190 }
191
192 return ::testing::AssertionSuccess();
193}
194
195template <typename ctype>
196::testing::AssertionResult assert_tensor_eq_with_dtype(

Callers

nothing calls this directly

Calls 3

diffFunction · 0.70
good_floatFunction · 0.70
to_stringMethod · 0.45

Tested by

no test coverage detected