| 324 | } |
| 325 | |
| 326 | xla::StatusOr<std::shared_ptr<XrtBuffer>> XrtExecutable::Execute( |
| 327 | const std::vector<std::shared_ptr<XrtBuffer>>& args) { |
| 328 | TF_RET_CHECK(device_assignment_.replica_count() == 1 && |
| 329 | device_assignment_.computation_count() == 1) |
| 330 | << device_assignment_.ToString(); |
| 331 | int xrt_device_ordinal = device_assignment_(0, 0); |
| 332 | int tf_device_id = context_->tf_device_ids().at(xrt_device_ordinal); |
| 333 | |
| 334 | TensorProto config_proto; |
| 335 | config_proto.set_dtype(DT_STRING); |
| 336 | config_proto.add_string_val(); |
| 337 | XrtTensorHandle execution_config_handle = |
| 338 | EnqueueConst(handle_.context().get(), tf_device_id, config_proto, |
| 339 | /*host_memory=*/true); |
| 340 | |
| 341 | protobuf::Map<string, AttrValue> attrs; |
| 342 | attrs["Ninputs"] = MakeAttrValue(args.size()); |
| 343 | |
| 344 | std::vector<const XrtTensorHandle*> inputs; |
| 345 | inputs.reserve(args.size() + 2); |
| 346 | inputs.push_back(&handle_); |
| 347 | inputs.push_back(&execution_config_handle); |
| 348 | for (const std::shared_ptr<XrtBuffer>& arg : args) { |
| 349 | if (arg->handle().device_id() != tf_device_id) { |
| 350 | return errors::InvalidArgument( |
| 351 | "Input buffer to Execute() is not on the device for which the " |
| 352 | "computation was compiled. Target device is ", |
| 353 | tf_device_id, ", buffer is on device ", arg->handle().device_id()); |
| 354 | } |
| 355 | inputs.push_back(&arg->handle()); |
| 356 | } |
| 357 | |
| 358 | XrtTensorHandle result_handle = std::move(handle_.context()->EnqueueOp( |
| 359 | "XRTExecute", inputs, /*output_arity=*/1, attrs, tf_device_id)[0]); |
| 360 | |
| 361 | return std::make_shared<XrtBuffer>(std::move(result_handle), |
| 362 | xrt_device_ordinal, shape_.result()); |
| 363 | } |
| 364 | |
| 365 | xla::StatusOr<xla::Array2D<std::shared_ptr<XrtBuffer>>> |
| 366 | XrtExecutable::ExecuteReplicated( |