| 413 | } |
| 414 | |
| 415 | void CheckTFE_TensorHandleHasFloats(TFE_TensorHandle* handle, |
| 416 | const std::vector<float>& expected_values) { |
| 417 | std::unique_ptr<TF_Status, decltype(&TF_DeleteStatus)> status( |
| 418 | TF_NewStatus(), TF_DeleteStatus); |
| 419 | TF_Tensor* t = TFE_TensorHandleResolve(handle, status.get()); |
| 420 | ASSERT_EQ(TF_OK, TF_GetCode(status.get())) << TF_Message(status.get()); |
| 421 | std::unique_ptr<float[]> actual_values(new float[expected_values.size()]); |
| 422 | EXPECT_EQ(sizeof(float) * expected_values.size(), TF_TensorByteSize(t)); |
| 423 | memcpy(actual_values.get(), TF_TensorData(t), TF_TensorByteSize(t)); |
| 424 | TF_DeleteTensor(t); |
| 425 | |
| 426 | for (int i = 0; i < expected_values.size(); i++) { |
| 427 | EXPECT_EQ(expected_values[i], actual_values[i]) |
| 428 | << "Mismatch in expected values at (zero-based) index " << i; |
| 429 | } |
| 430 | } |
| 431 | |
| 432 | void CheckRemoteMatMulExecutesOK(TFE_Context* ctx, |
| 433 | const char* remote_device_name, |
no test coverage detected