| 58 | }; |
| 59 | |
| 60 | TEST(DeviceUtilTest, AddDeviceToOp) { |
| 61 | mlir::MLIRContext context; |
| 62 | mlir::OwningModuleRef module_ref = |
| 63 | mlir::ModuleOp::create(mlir::UnknownLoc::get(&context)); |
| 64 | |
| 65 | DeviceSet device_set; |
| 66 | llvm::SmallVector<std::unique_ptr<Device>, 2> devices; |
| 67 | devices.push_back( |
| 68 | FakeDevice::Make("/job:worker/replica:0/task:0/device:CPU:0")); |
| 69 | devices.push_back( |
| 70 | FakeDevice::Make("/job:worker/replica:1/task:2/device:GPU:3")); |
| 71 | for (auto& device : devices) device_set.AddDevice(device.get()); |
| 72 | |
| 73 | AddDevicesToOp(*module_ref, &device_set); |
| 74 | auto devices_attr = module_ref->getAttrOfType<mlir::ArrayAttr>("tf.devices"); |
| 75 | ASSERT_NE(devices_attr, nullptr); |
| 76 | ASSERT_EQ(devices_attr.size(), 2); |
| 77 | auto device_attr_0 = devices_attr.getValue()[0].dyn_cast<mlir::StringAttr>(); |
| 78 | ASSERT_NE(device_attr_0, nullptr); |
| 79 | EXPECT_EQ(device_attr_0.getValue(), |
| 80 | "/job:worker/replica:0/task:0/device:CPU:0"); |
| 81 | auto device_attr_1 = devices_attr.getValue()[1].dyn_cast<mlir::StringAttr>(); |
| 82 | ASSERT_NE(device_attr_1, nullptr); |
| 83 | EXPECT_EQ(device_attr_1.getValue(), |
| 84 | "/job:worker/replica:1/task:2/device:GPU:3"); |
| 85 | } |
| 86 | |
| 87 | TEST(DeviceUtilTest, AddDeviceToOpNullDeviceSet) { |
| 88 | mlir::MLIRContext context; |
nothing calls this directly
no test coverage detected