| 1321 | } |
| 1322 | |
| 1323 | XlaOp XlaBuilder::Infeed(const Shape& shape, const string& config) { |
| 1324 | return ReportErrorOrReturn([&]() -> StatusOr<XlaOp> { |
| 1325 | HloInstructionProto instr; |
| 1326 | if (!LayoutUtil::HasLayout(shape)) { |
| 1327 | return InvalidArgument("Given shape to Infeed must have a layout"); |
| 1328 | } |
| 1329 | const Shape infeed_instruction_shape = |
| 1330 | ShapeUtil::MakeTupleShape({shape, ShapeUtil::MakeTokenShape()}); |
| 1331 | *instr.mutable_shape() = infeed_instruction_shape.ToProto(); |
| 1332 | instr.set_infeed_config(config); |
| 1333 | |
| 1334 | if (shape.IsArray() && sharding() && |
| 1335 | sharding()->type() == OpSharding::OTHER) { |
| 1336 | // TODO(b/110793772): Support tiled array-shaped infeeds. |
| 1337 | return InvalidArgument( |
| 1338 | "Tiled sharding is not yet supported for array-shaped infeeds"); |
| 1339 | } |
| 1340 | |
| 1341 | if (sharding() && sharding()->type() == OpSharding::REPLICATED) { |
| 1342 | return InvalidArgument( |
| 1343 | "Replicated sharding is not yet supported for infeeds"); |
| 1344 | } |
| 1345 | |
| 1346 | // Infeed takes a single token operand. Generate the token to pass to the |
| 1347 | // infeed. |
| 1348 | XlaOp token; |
| 1349 | auto make_token = [&]() { |
| 1350 | HloInstructionProto token_instr; |
| 1351 | *token_instr.mutable_shape() = ShapeUtil::MakeTokenShape().ToProto(); |
| 1352 | return AddInstruction(std::move(token_instr), HloOpcode::kAfterAll, {}); |
| 1353 | }; |
| 1354 | if (sharding()) { |
| 1355 | // Arbitrarily assign token to device 0. |
| 1356 | OpSharding sharding = sharding_builder::AssignDevice(0); |
| 1357 | XlaScopedShardingAssignment scoped_sharding(this, sharding); |
| 1358 | TF_ASSIGN_OR_RETURN(token, make_token()); |
| 1359 | } else { |
| 1360 | TF_ASSIGN_OR_RETURN(token, make_token()); |
| 1361 | } |
| 1362 | |
| 1363 | // The sharding is set by the client according to the data tuple shape. |
| 1364 | // However, the shape of the infeed instruction is a tuple containing the |
| 1365 | // data and a token. For tuple sharding type, the sharding must be changed |
| 1366 | // to accommodate the token. |
| 1367 | XlaOp infeed; |
| 1368 | if (sharding() && sharding()->type() == OpSharding::TUPLE) { |
| 1369 | // TODO(b/80000000): Remove this when clients have been updated to handle |
| 1370 | // tokens. |
| 1371 | OpSharding infeed_instruction_sharding = *sharding(); |
| 1372 | // Arbitrarily assign the token to device 0. |
| 1373 | *infeed_instruction_sharding.add_tuple_shardings() = |
| 1374 | sharding_builder::AssignDevice(0); |
| 1375 | XlaScopedShardingAssignment scoped_sharding(this, |
| 1376 | infeed_instruction_sharding); |
| 1377 | TF_ASSIGN_OR_RETURN(infeed, AddInstruction(std::move(instr), |
| 1378 | HloOpcode::kInfeed, {token})); |
| 1379 | } else { |
| 1380 | TF_ASSIGN_OR_RETURN(infeed, AddInstruction(std::move(instr), |
no test coverage detected