| 54 | } |
| 55 | |
| 56 | TEST(SelectedRows, SparseTable) { |
| 57 | phi::CPUPlace cpu; |
| 58 | SelectedRows table; |
| 59 | |
| 60 | int64_t table_size = 100; |
| 61 | int64_t embedding_width = 8; |
| 62 | // initialize a sparse table |
| 63 | table.mutable_value()->Resize( |
| 64 | common::make_ddim({table_size, embedding_width})); |
| 65 | auto* data = table.mutable_value()->mutable_data<float>(cpu); |
| 66 | for (int64_t i = 0; i < table_size; ++i) { |
| 67 | for (int64_t j = 0; j < embedding_width; ++j) { |
| 68 | data[i * embedding_width + j] = static_cast<float>(i); |
| 69 | } |
| 70 | } |
| 71 | ASSERT_EQ(table.AutoGrownIndex(10, true, false), 0); |
| 72 | ASSERT_EQ(table.AutoGrownIndex(8, true, false), 1); |
| 73 | ASSERT_EQ(table.AutoGrownIndex(8, true, false), 1); |
| 74 | ASSERT_EQ(table.AutoGrownIndex(6, true, false), 2); |
| 75 | for (int64_t i = 11; i < 20; i++) { |
| 76 | ASSERT_EQ(table.AutoGrownIndex(i, true, true), -1); |
| 77 | ASSERT_TRUE(!table.HasKey(i)); |
| 78 | } |
| 79 | ASSERT_TRUE(table.HasKey(10)); |
| 80 | ASSERT_TRUE(table.HasKey(8)); |
| 81 | ASSERT_TRUE(table.HasKey(6)); |
| 82 | ASSERT_EQ(table.rows().size(), 3UL); |
| 83 | |
| 84 | phi::DenseTensor ids; |
| 85 | ids.Resize(common::make_ddim({4})); |
| 86 | auto* ids_data = ids.mutable_data<int64_t>(cpu); |
| 87 | ids_data[0] = static_cast<int64_t>(6); |
| 88 | ids_data[1] = static_cast<int64_t>(6); |
| 89 | ids_data[2] = static_cast<int64_t>(8); |
| 90 | ids_data[3] = static_cast<int64_t>(10); |
| 91 | |
| 92 | phi::DenseTensor get_value; |
| 93 | auto* value_data = get_value.mutable_data<float>( |
| 94 | common::make_ddim({4, embedding_width}), cpu); |
| 95 | table.Get(ids, &get_value); |
| 96 | |
| 97 | for (int j = 0; j < embedding_width; ++j) { |
| 98 | ASSERT_EQ(value_data[0 * embedding_width + j], 2); |
| 99 | } |
| 100 | for (int j = 0; j < embedding_width; ++j) { |
| 101 | ASSERT_EQ(value_data[1 * embedding_width + j], 2); |
| 102 | } |
| 103 | for (int j = 0; j < embedding_width; ++j) { |
| 104 | ASSERT_EQ(value_data[2 * embedding_width + j], 1); |
| 105 | } |
| 106 | for (int j = 0; j < embedding_width; ++j) { |
| 107 | ASSERT_EQ(value_data[3 * embedding_width + j], 0); |
| 108 | } |
| 109 | } |
| 110 | |
| 111 | void f1(SelectedRows* table, int table_size) { |
| 112 | for (int i = 1000000; i > 0; --i) { |
nothing calls this directly
no test coverage detected