| 3365 | } |
| 3366 | |
| 3367 | std::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()); |
nothing calls this directly
no test coverage detected