| 36 | public ::testing::WithParamInterface<std::tuple<int, int, int, int>> {}; |
| 37 | |
| 38 | XLA_TEST_P(TridiagonalTest, Solves) { |
| 39 | const auto& spec = GetParam(); |
| 40 | xla::XlaBuilder builder(TestName()); |
| 41 | |
| 42 | const int64 num_eqs = 5; |
| 43 | const int64 num_rhs = 3; |
| 44 | const int64 lower_diagonal_batch_size = std::get<0>(spec); |
| 45 | const int64 main_diagonal_batch_size = std::get<1>(spec); |
| 46 | const int64 upper_diagonal_batch_size = std::get<2>(spec); |
| 47 | const int64 rhs_diagonal_batch_size = std::get<2>(spec); |
| 48 | |
| 49 | const int64 max_batch_size = |
| 50 | std::max({lower_diagonal_batch_size, main_diagonal_batch_size, |
| 51 | upper_diagonal_batch_size, rhs_diagonal_batch_size}); |
| 52 | |
| 53 | Array3D<float> lower_diagonal(lower_diagonal_batch_size, 1, num_eqs); |
| 54 | Array3D<float> main_diagonal(main_diagonal_batch_size, 1, num_eqs); |
| 55 | Array3D<float> upper_diagonal(upper_diagonal_batch_size, 1, num_eqs); |
| 56 | Array3D<float> rhs(rhs_diagonal_batch_size, num_rhs, num_eqs); |
| 57 | |
| 58 | lower_diagonal.FillRandom(1.0, /*mean=*/0.0, /*seed=*/0); |
| 59 | main_diagonal.FillRandom(0.05, /*mean=*/1.0, |
| 60 | /*seed=*/max_batch_size * num_eqs); |
| 61 | upper_diagonal.FillRandom(1.0, /*mean=*/0.0, |
| 62 | /*seed=*/2 * max_batch_size * num_eqs); |
| 63 | rhs.FillRandom(1.0, /*mean=*/0.0, /*seed=*/3 * max_batch_size * num_eqs); |
| 64 | |
| 65 | XlaOp lower_diagonal_xla; |
| 66 | XlaOp main_diagonal_xla; |
| 67 | XlaOp upper_diagonal_xla; |
| 68 | XlaOp rhs_xla; |
| 69 | |
| 70 | auto lower_diagonal_data = CreateR3Parameter<float>( |
| 71 | lower_diagonal, 0, "lower_diagonal", &builder, &lower_diagonal_xla); |
| 72 | auto main_diagonal_data = CreateR3Parameter<float>( |
| 73 | main_diagonal, 1, "main_diagonal", &builder, &main_diagonal_xla); |
| 74 | auto upper_diagonal_data = CreateR3Parameter<float>( |
| 75 | upper_diagonal, 2, "upper_diagonal", &builder, &upper_diagonal_xla); |
| 76 | auto rhs_data = CreateR3Parameter<float>(rhs, 3, "rhs", &builder, &rhs_xla); |
| 77 | |
| 78 | TF_ASSERT_OK_AND_ASSIGN(XlaOp x, |
| 79 | ThomasSolver(lower_diagonal_xla, main_diagonal_xla, |
| 80 | upper_diagonal_xla, rhs_xla)); |
| 81 | |
| 82 | auto Coefficient = [](auto operand, auto i) { |
| 83 | return SliceInMinorDims(operand, /*start=*/{i}, /*end=*/{i + 1}); |
| 84 | }; |
| 85 | |
| 86 | std::vector<XlaOp> relative_errors(num_eqs); |
| 87 | |
| 88 | for (int64 i = 0; i < num_eqs; i++) { |
| 89 | auto a_i = Coefficient(lower_diagonal_xla, i); |
| 90 | auto b_i = Coefficient(main_diagonal_xla, i); |
| 91 | auto c_i = Coefficient(upper_diagonal_xla, i); |
| 92 | auto d_i = Coefficient(rhs_xla, i); |
| 93 | |
| 94 | if (i == 0) { |
| 95 | relative_errors[i] = |
nothing calls this directly
no test coverage detected