A simple test on writing a few tensor slices TODO(yangke): refactor into smaller tests: will do as we add more stuff to the writer.
| 71 | // TODO(yangke): refactor into smaller tests: will do as we add more stuff to |
| 72 | // the writer. |
| 73 | TEST(TensorSliceWriteTest, SimpleWrite) { |
| 74 | const string filename = io::JoinPath(testing::TmpDir(), "checkpoint"); |
| 75 | |
| 76 | TensorSliceWriter writer(filename, CreateTableTensorSliceBuilder); |
| 77 | |
| 78 | // Add some int32 tensor slices |
| 79 | { |
| 80 | TensorShape shape({5, 10}); |
| 81 | TensorSlice slice = TensorSlice::ParseOrDie("-:0,1"); |
| 82 | const int32 data[] = {0, 1, 2, 3, 4}; |
| 83 | TF_CHECK_OK(writer.Add("test", shape, slice, data)); |
| 84 | } |
| 85 | |
| 86 | // Two slices share the same tensor name |
| 87 | { |
| 88 | TensorShape shape({5, 10}); |
| 89 | TensorSlice slice = TensorSlice::ParseOrDie("-:3,1"); |
| 90 | const int32 data[] = {10, 11, 12, 13, 14}; |
| 91 | TF_CHECK_OK(writer.Add("test", shape, slice, data)); |
| 92 | } |
| 93 | |
| 94 | // Another slice from a different float tensor -- it has a different name and |
| 95 | // should be inserted in front of the previous tensor |
| 96 | { |
| 97 | TensorShape shape({3, 2}); |
| 98 | TensorSlice slice = TensorSlice::ParseOrDie("-:-"); |
| 99 | const float data[] = {1.2, 1.3, 1.4, 2.1, 2.2, 2.3}; |
| 100 | TF_CHECK_OK(writer.Add("AA", shape, slice, data)); |
| 101 | } |
| 102 | |
| 103 | // A slice with int64 data |
| 104 | { |
| 105 | TensorShape shape({5, 10}); |
| 106 | TensorSlice slice = TensorSlice::ParseOrDie("-:3,1"); |
| 107 | const int64 data[] = {10, 11, 12, 13, 14}; |
| 108 | TF_CHECK_OK(writer.Add("int64", shape, slice, data)); |
| 109 | } |
| 110 | |
| 111 | // A slice with int16 data |
| 112 | { |
| 113 | TensorShape shape({5, 10}); |
| 114 | TensorSlice slice = TensorSlice::ParseOrDie("-:3,1"); |
| 115 | const int16 data[] = {10, 11, 12, 13, 14}; |
| 116 | TF_CHECK_OK(writer.Add("int16", shape, slice, data)); |
| 117 | } |
| 118 | |
| 119 | TF_CHECK_OK(writer.Finish()); |
| 120 | |
| 121 | // Now we examine the checkpoint file manually. |
| 122 | TensorSliceWriteTestHelper::CheckEntries(filename); |
| 123 | } |
| 124 | |
| 125 | } // namespace |
| 126 |