MCPcopy Create free account
hub / github.com/DeepRec-AI/DeepRec / XLA_TEST_P

Function XLA_TEST_P

tensorflow/compiler/xla/client/lib/tridiagonal_test.cc:38–119  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

36 public ::testing::WithParamInterface<std::tuple<int, int, int, int>> {};
37
38XLA_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] =

Callers

nothing calls this directly

Calls 8

GetParamFunction · 0.85
TestNameFunction · 0.85
SliceInMinorDimsFunction · 0.85
CoefficientFunction · 0.85
ConcatInDimFunction · 0.85
maxFunction · 0.50
AbsFunction · 0.50
FillRandomMethod · 0.45

Tested by

no test coverage detected