| 34 | } |
| 35 | |
| 36 | void require_close( |
| 37 | const std::vector<float> & actual, |
| 38 | const engine::assets::TensorDataF32 & expected, |
| 39 | float max_allowed, |
| 40 | double mean_allowed, |
| 41 | const std::string & label) { |
| 42 | require(expected.shape.rank == 2 && expected.shape.dims[0] == 1, label + " expected shape mismatch"); |
| 43 | require(actual.size() == expected.values.size(), label + " size mismatch"); |
| 44 | float max_diff = 0.0f; |
| 45 | size_t max_index = 0; |
| 46 | double mean_diff = 0.0; |
| 47 | for (size_t i = 0; i < actual.size(); ++i) { |
| 48 | const float diff = std::fabs(actual[i] - expected.values[i]); |
| 49 | mean_diff += static_cast<double>(diff); |
| 50 | if (diff > max_diff) { |
| 51 | max_diff = diff; |
| 52 | max_index = i; |
| 53 | } |
| 54 | } |
| 55 | mean_diff /= static_cast<double>(actual.size()); |
| 56 | if (max_diff > max_allowed || mean_diff > mean_allowed) { |
| 57 | std::ostringstream oss; |
| 58 | oss << label << " mismatch: max_diff=" << max_diff |
| 59 | << " mean_diff=" << mean_diff |
| 60 | << " index=" << max_index |
| 61 | << " expected=" << expected.values[max_index] |
| 62 | << " actual=" << actual[max_index]; |
| 63 | throw std::runtime_error(oss.str()); |
| 64 | } |
| 65 | } |
| 66 | |
| 67 | void run_case(int case_index) { |
| 68 | const auto model = engine::audio::FlashSrModel::load_from_directory( |