| 2486 | } |
| 2487 | |
| 2488 | XlaOp XlaBuilder::SendToHost(XlaOp operand, XlaOp token, |
| 2489 | const Shape& shape_with_layout, |
| 2490 | const ChannelHandle& handle) { |
| 2491 | return ReportErrorOrReturn([&]() -> StatusOr<XlaOp> { |
| 2492 | if (!LayoutUtil::HasLayout(shape_with_layout)) { |
| 2493 | return InvalidArgument("Shape passed to SendToHost must have a layout"); |
| 2494 | } |
| 2495 | TF_ASSIGN_OR_RETURN(const Shape* operand_shape, GetShapePtr(operand)); |
| 2496 | if (!ShapeUtil::Compatible(*operand_shape, shape_with_layout)) { |
| 2497 | return InvalidArgument( |
| 2498 | "SendToHost shape %s must be compatible with operand shape %s", |
| 2499 | ShapeUtil::HumanStringWithLayout(shape_with_layout), |
| 2500 | ShapeUtil::HumanStringWithLayout(*operand_shape)); |
| 2501 | } |
| 2502 | // TODO(b/111544877): Support tuple shapes. |
| 2503 | if (!operand_shape->IsArray()) { |
| 2504 | return InvalidArgument("SendToHost only supports array shapes, shape: %s", |
| 2505 | ShapeUtil::HumanString(*operand_shape)); |
| 2506 | } |
| 2507 | |
| 2508 | if (handle.type() != ChannelHandle::DEVICE_TO_HOST) { |
| 2509 | return InvalidArgument("SendToHost must use a device-to-host channel"); |
| 2510 | } |
| 2511 | |
| 2512 | // Send instruction produces a tuple of {aliased operand, U32 context, |
| 2513 | // token}. |
| 2514 | HloInstructionProto send_instr; |
| 2515 | *send_instr.mutable_shape() = |
| 2516 | ShapeUtil::MakeTupleShape({shape_with_layout, |
| 2517 | ShapeUtil::MakeShape(U32, {}), |
| 2518 | ShapeUtil::MakeTokenShape()}) |
| 2519 | .ToProto(); |
| 2520 | send_instr.set_channel_id(handle.handle()); |
| 2521 | send_instr.set_is_host_transfer(true); |
| 2522 | TF_ASSIGN_OR_RETURN(XlaOp send, |
| 2523 | AddInstruction(std::move(send_instr), HloOpcode::kSend, |
| 2524 | {operand, token})); |
| 2525 | |
| 2526 | HloInstructionProto send_done_instr; |
| 2527 | *send_done_instr.mutable_shape() = ShapeUtil::MakeTokenShape().ToProto(); |
| 2528 | send_done_instr.set_channel_id(handle.handle()); |
| 2529 | send_done_instr.set_is_host_transfer(true); |
| 2530 | return AddInstruction(std::move(send_done_instr), HloOpcode::kSendDone, |
| 2531 | {send}); |
| 2532 | }); |
| 2533 | } |
| 2534 | |
| 2535 | XlaOp XlaBuilder::RecvFromHost(XlaOp token, const Shape& shape, |
| 2536 | const ChannelHandle& handle) { |
no test coverage detected