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

Method EmitThreadLocalCall

tensorflow/compiler/xla/service/cpu/ir_emitter.cc:3367–3429  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

3365}
3366
3367std::vector<llvm::Value*> IrEmitter::EmitThreadLocalCall(
3368 const HloComputation& callee, absl::Span<llvm::Value* const> parameters,
3369 absl::string_view name) {
3370 CHECK(absl::c_binary_search(thread_local_computations_, &callee));
3371 const Shape& return_shape = callee.root_instruction()->shape();
3372 bool is_scalar_return = ShapeUtil::IsScalar(return_shape);
3373 bool is_tuple_of_scalars_return =
3374 return_shape.IsTuple() &&
3375 absl::c_all_of(return_shape.tuple_shapes(), [&](const Shape& shape) {
3376 return ShapeUtil::IsScalar(shape);
3377 });
3378 CHECK(is_scalar_return || is_tuple_of_scalars_return);
3379
3380 std::vector<llvm::Value*> parameter_addrs;
3381 for (llvm::Value* parameter : parameters) {
3382 CHECK(!parameter->getType()->isPointerTy());
3383 llvm::Value* parameter_addr = llvm_ir::EmitAllocaAtFunctionEntry(
3384 parameter->getType(), "arg_addr", &b_);
3385 Store(parameter, parameter_addr);
3386 parameter_addrs.push_back(parameter_addr);
3387 }
3388
3389 llvm::Type* return_value_buffer_type =
3390 llvm_ir::ShapeToIrType(return_shape, module_);
3391 std::string retval_alloca_name = absl::StrCat(name, "_return_value_addr");
3392 int retval_alignment =
3393 is_scalar_return
3394 ? MinimumAlignmentForPrimitiveType(return_shape.element_type())
3395 : 0;
3396 llvm::Value* return_value_buffer = llvm_ir::EmitAllocaAtFunctionEntry(
3397 return_value_buffer_type, retval_alloca_name, &b_, retval_alignment);
3398
3399 std::vector<llvm::Value*> allocas_for_returned_scalars;
3400 if (is_scalar_return) {
3401 allocas_for_returned_scalars.push_back(return_value_buffer);
3402 } else {
3403 constexpr int max_tuple_size = 1000;
3404 CHECK_LT(return_shape.tuple_shapes_size(), max_tuple_size)
3405 << "Multivalue function can not return more than 1000 elements to avoid"
3406 << " stack smashing";
3407 allocas_for_returned_scalars =
3408 llvm_ir::EmitTupleAllocasAtFunctionEntry(return_shape, &b_);
3409 llvm_ir::IrArray tuple_array(return_value_buffer, return_shape);
3410
3411 EmitTuple(tuple_array, allocas_for_returned_scalars, &b_);
3412 }
3413
3414 Call(FindOrDie(emitted_functions_, &callee),
3415 GetArrayFunctionCallArguments(
3416 parameter_addrs, &b_, name,
3417 /*return_value_buffer=*/return_value_buffer,
3418 /*exec_run_options_arg=*/GetExecutableRunOptionsArgument(),
3419 /*buffer_table_arg=*/
3420 llvm::Constant::getNullValue(b_.getInt8PtrTy()->getPointerTo()),
3421 /*profile_counters_arg=*/GetProfileCountersArgument()));
3422
3423 std::vector<llvm::Value*> returned_scalars;
3424 returned_scalars.reserve(allocas_for_returned_scalars.size());

Callers

nothing calls this directly

Calls 15

ShapeToIrTypeFunction · 0.85
EmitTupleFunction · 0.85
LoadFunction · 0.85
root_instructionMethod · 0.80
tuple_shapes_sizeMethod · 0.80
IsScalarFunction · 0.50
StrCatFunction · 0.50
CallFunction · 0.50
shapeMethod · 0.45

Tested by

no test coverage detected