MCPcopy Create free account
hub / github.com/DeepRec-AI/DeepRec / Infeed

Method Infeed

tensorflow/compiler/xla/client/xla_builder.cc:1323–1394  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

1321}
1322
1323XlaOp 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),

Callers 1

InfeedFunction · 0.45

Calls 9

InvalidArgumentFunction · 0.85
AssignDeviceFunction · 0.85
typeMethod · 0.65
TF_ASSIGN_OR_RETURNFunction · 0.50
mutable_shapeMethod · 0.45
ToProtoMethod · 0.45
set_infeed_configMethod · 0.45
IsArrayMethod · 0.45
set_tuple_indexMethod · 0.45

Tested by

no test coverage detected