PiperOrigin-RevId: 987427673
This commit is contained in:
Enver Kayaaslan
2026-09-25 08:33:07 -07:00
committed by TensorFlower Gardener
parent 8c9cfc9f5a
commit d9ecfe9451
4 changed files with 593 additions and 14 deletions
+3 -1
View File
@@ -1217,7 +1217,6 @@ cc_library(
"@com_google_absl//absl/strings:str_format",
"@com_google_absl//absl/strings:string_view",
"@com_google_absl//absl/time",
"@tsl//tsl/platform:errors",
],
)
@@ -1225,6 +1224,7 @@ xla_cc_test(
name = "hlo_constant_folding_test",
srcs = ["hlo_constant_folding_test.cc"],
deps = [
":flatten_call_graph",
":hlo_constant_folding",
"//xla:literal",
"//xla:literal_util",
@@ -1233,12 +1233,14 @@ xla_cc_test(
"//xla:xla_data_proto_cc",
"//xla/hlo/ir:hlo",
"//xla/hlo/parser:hlo_parser",
"//xla/hlo/testlib:filecheck",
"//xla/hlo/testlib:hlo_hardware_independent_test_base",
"//xla/hlo/testlib:pattern_matcher_gmock",
"//xla/hlo/testlib:test",
"//xla/hlo/utils:hlo_matchers",
"//xla/service:pattern_matcher",
"//xla/tsl/platform:statusor",
"@com_google_absl//absl/status:status_matchers",
"@com_google_absl//absl/strings:string_view",
"@com_google_absl//absl/types:span",
"@com_google_googletest//:gtest_main",
@@ -19,6 +19,7 @@ limitations under the License.
#include <atomic>
#include <cstdint>
#include <memory>
#include <optional>
#include <utility>
#include <vector>
@@ -47,7 +48,6 @@ limitations under the License.
#include "xla/tsl/platform/errors.h"
#include "xla/tsl/platform/statusor.h"
#include "xla/xla_data.pb.h"
#include "tsl/platform/errors.h"
namespace xla {
@@ -447,6 +447,84 @@ absl::StatusOr<bool> PropagateIdenticalConstantArguments(
return changed;
}
// TODO(enver): Share CloneComputation with other passes, such as
// FlattenCallGraph, by moving to HloModule or HloComputation.
HloComputation* CloneComputation(HloModule* module,
HloComputation* computation) {
if (!module->has_schedule() ||
!module->schedule().is_computation_scheduled(computation)) {
return module->AddEmbeddedComputation(computation->Clone());
}
auto [clone, clone_sequence] = computation->CloneWithSchedule();
HloComputation* clone_ptr = module->AddEmbeddedComputation(std::move(clone));
module->schedule().set_sequence(clone_ptr, clone_sequence);
return clone_ptr;
}
using SpecializationKey = HloConstantFolding::SpecializationKey;
absl::StatusOr<bool> SpecializeCalls(
HloModule* module, HloComputation* computation,
absl::flat_hash_map<SpecializationKey, HloComputation*>&
specialization_cache,
std::vector<HloComputation*>& computation_versions) {
auto caller_instructions = computation->caller_instructions();
if (caller_instructions.empty()) {
return false;
}
if (!absl::c_all_of(caller_instructions, [](const HloInstruction* instr) {
return instr->opcode() == HloOpcode::kCall;
})) {
return false;
}
if (caller_instructions.size() > 1) {
// Sort the caller instructions by their unique id to make the compilation
// deterministic.
absl::c_sort(caller_instructions,
[](const HloInstruction* a, const HloInstruction* b) {
return a->unique_id() < b->unique_id();
});
}
bool changed = false;
bool original_computation_used = false;
for (HloInstruction* caller : caller_instructions) {
std::vector<std::optional<LiteralSlice>> arguments;
arguments.reserve(caller->operand_count());
for (const HloInstruction* operand : caller->operands()) {
if (operand->opcode() == HloOpcode::kConstant) {
arguments.push_back(LiteralSlice(operand->literal()));
} else {
arguments.push_back(std::nullopt);
}
}
SpecializationKey key{computation, std::move(arguments)};
auto it = specialization_cache.find(key);
HloComputation* target_comp = nullptr;
if (it != specialization_cache.end()) {
target_comp = it->second;
} else {
if (!original_computation_used) {
// The first argument combination for this computation uses the original
// computation.
target_comp = computation;
original_computation_used = true;
} else {
target_comp = CloneComputation(module, computation);
computation_versions.push_back(target_comp);
}
specialization_cache.emplace(key, target_comp);
}
if (caller->to_apply() != target_comp) {
caller->set_to_apply(target_comp);
changed = true;
}
}
return changed;
}
} // namespace
absl::StatusOr<bool> HloConstantFolding::RunOnComputation(
@@ -606,6 +684,7 @@ absl::StatusOr<bool> HloConstantFolding::RunOnComputation(
absl::StatusOr<bool> HloConstantFolding::RunImpl(
HloModule* module,
const absl::flat_hash_set<absl::string_view>& execution_threads) {
specialization_cache_.clear();
// Limit the constant folding to 0 iterations to skip folding loops in the
// default case. This retains the behavior from before while loop support in
// HloEvaluator and may be revised.
@@ -627,10 +706,21 @@ absl::StatusOr<bool> HloConstantFolding::RunImpl(
// arguments from callers to callees.
for (auto it = computations.rbegin(); it != computations.rend(); ++it) {
HloComputation* computation = *it;
ABSL_ASSIGN_OR_RETURN(bool computation_changed,
RunOnComputation(computation, evaluator.get(),
is_foldable_computation));
changed |= computation_changed;
// TODO(b/260601110): Early exit the computation if all callers are already
// constant folded.
std::vector<HloComputation*> computation_versions;
computation_versions.push_back(computation);
ABSL_ASSIGN_OR_RETURN(bool did_specialize,
SpecializeCalls(module, computation, specialization_cache_,
computation_versions));
changed |= did_specialize;
for (HloComputation* computation_version : computation_versions) {
ABSL_ASSIGN_OR_RETURN(bool version_changed,
RunOnComputation(computation_version, evaluator.get(),
is_foldable_computation));
changed |= version_changed;
}
}
return changed;
}
@@ -19,13 +19,16 @@ limitations under the License.
#include <atomic>
#include <cstdint>
#include <functional>
#include <optional>
#include <vector>
#include "absl/container/flat_hash_map.h"
#include "absl/container/flat_hash_set.h"
#include "absl/status/statusor.h"
#include "absl/strings/string_view.h"
#include "xla/hlo/ir/hlo_module.h"
#include "xla/hlo/ir/hlo_computation.h"
#include "xla/hlo/pass/hlo_pass_interface.h"
#include "xla/literal.h"
#include "xla/shape.h"
namespace xla {
@@ -82,6 +85,49 @@ class HloConstantFolding : public HloModulePass {
explicit HloConstantFolding(const Options& options) : options_(options) {}
absl::string_view name() const override { return "constant_folding"; }
struct SpecializationKey {
const HloComputation* original_computation = nullptr;
std::vector<std::optional<LiteralSlice>> arguments;
bool operator==(const SpecializationKey& other) const {
if (original_computation != other.original_computation ||
arguments.size() != other.arguments.size()) {
return false;
}
for (size_t i = 0; i < arguments.size(); ++i) {
if (arguments[i].has_value() != other.arguments[i].has_value()) {
return false;
}
if (arguments[i].has_value()) {
if (!arguments[i]->Equal(*other.arguments[i],
/*layout_sensitive=*/true)) {
return false;
}
}
}
return true;
}
template <typename H>
friend H AbslHashValue(H h, const SpecializationKey& key) {
h = H::combine(std::move(h), key.original_computation,
key.arguments.size());
for (const auto& arg : key.arguments) {
if (arg.has_value()) {
h = H::combine(std::move(h), true, Literal::AbslHashable<true>(*arg));
} else {
h = H::combine(std::move(h), false);
}
}
return h;
}
};
const absl::flat_hash_map<SpecializationKey, HloComputation*>&
specialization_cache() const {
return specialization_cache_;
}
protected:
// Run constant folding operations on the given module. Returns whether the
// module was changed (constant expressions folded).
@@ -99,6 +145,7 @@ class HloConstantFolding : public HloModulePass {
static std::atomic<int64_t> slow_op_counter_;
Options options_;
absl::flat_hash_map<SpecializationKey, HloComputation*> specialization_cache_;
};
} // namespace xla
@@ -21,15 +21,18 @@ limitations under the License.
#include <utility>
#include <vector>
#include "absl/status/status_matchers.h"
#include "absl/strings/string_view.h"
#include "absl/types/span.h"
#include "xla/hlo/ir/hlo_computation.h"
#include "xla/hlo/ir/hlo_instruction.h"
#include "xla/hlo/ir/hlo_opcode.h"
#include "xla/hlo/parser/hlo_parser.h"
#include "xla/hlo/testlib/filecheck.h"
#include "xla/hlo/testlib/hlo_hardware_independent_test_base.h"
#include "xla/hlo/testlib/pattern_matcher_gmock.h"
#include "xla/hlo/testlib/test.h"
#include "xla/hlo/transforms/simplifiers/flatten_call_graph.h"
#include "xla/hlo/utils/hlo_matchers.h"
#include "xla/layout_util.h"
#include "xla/literal.h"
@@ -46,6 +49,8 @@ limitations under the License.
namespace xla {
namespace {
using ::absl_testing::IsOkAndHolds;
namespace op = xla::testing::opcode_matchers;
namespace m = xla::match;
using HloConstantFoldingTest = HloHardwareIndependentTestBase;
@@ -759,7 +764,7 @@ TEST_F(HloConstantFoldingTest,
HloConstantFolding constant_folding;
TF_ASSERT_OK_AND_ASSIGN(bool result,
RunHloPass(&constant_folding, module.get()));
EXPECT_FALSE(result);
EXPECT_TRUE(result);
}
TEST_F(HloConstantFoldingTest, InterproceduralMultipleCallsitesSomeConstants) {
@@ -792,11 +797,10 @@ TEST_F(HloConstantFoldingTest, InterproceduralMultipleCallsitesSomeConstants) {
RunHloPass(&constant_folding, module.get()));
EXPECT_TRUE(result);
HloComputation* fn = module->GetComputationWithName("Fn");
EXPECT_THAT(fn->root_instruction(),
GmockMatch(m::Add(m::Subtract(m::Multiply(m::ConstantScalar(1),
m::Parameter(1)),
m::Parameter(2)),
m::ConstantScalar(2))));
EXPECT_THAT(
fn->root_instruction(),
GmockMatch(m::Add(m::Subtract(m::ConstantScalar(2), m::Parameter(2)),
m::ConstantScalar(2))));
}
TEST_F(HloConstantFoldingTest,
@@ -832,7 +836,7 @@ TEST_F(HloConstantFoldingTest,
EXPECT_TRUE(result);
HloComputation* fn = module->GetComputationWithName("Fn");
EXPECT_THAT(fn->root_instruction(),
GmockMatch(m::Subtract(m::Add(m::Add(), m::Parameter(1)),
GmockMatch(m::Subtract(m::Add(m::Constant(), m::Parameter(1)),
m::Constant())));
}
@@ -1400,5 +1404,441 @@ TEST_F(HloConstantFoldingTest, LateOptionsDontFoldGteWithControlDependency) {
EXPECT_FALSE(result);
}
TEST_F(HloConstantFoldingTest, ComputationCalledTwiceOnDifferentConstants) {
const std::string hlo_string = R"(
HloModule test
// CHECK-LABEL: %subcomp (
// CHECK: constant({15, 15})
// CHECK-LABEL: %subcomp.clone (
// CHECK: constant({25, 25})
subcomp {
p0 = s32[2] parameter(0)
p1 = s32[2] parameter(1)
c = s32[2] constant({5, 5})
add0 = s32[2] add(p0, c)
ROOT add1 = s32[2] add(p1, add0)
}
// CHECK-LABEL: ENTRY %entry (
// CHECK: %call1 = s32[2]{0} call(%c10, %p0.1), to_apply=%subcomp
// CHECK: %call2 = s32[2]{0} call(%c20, %p1.1), to_apply=%subcomp.clone
ENTRY entry {
p0.1 = s32[2] parameter(0)
p1.1 = s32[2] parameter(1)
c10 = s32[2] constant({10, 10})
c20 = s32[2] constant({20, 20})
call1 = s32[2] call(c10, p0.1), to_apply=subcomp
call2 = s32[2] call(c20, p1.1), to_apply=subcomp
ROOT out = s32[2] add(call1, call2)
})";
{
ASSERT_OK_AND_ASSIGN(auto module, ParseAndReturnVerifiedModule(hlo_string));
HloConstantFolding constant_folding;
EXPECT_THAT(constant_folding.Run(module.get()), IsOkAndHolds(true));
}
{
ASSERT_OK_AND_ASSIGN(auto module, ParseAndReturnVerifiedModule(hlo_string));
FlattenCallGraph flatten;
EXPECT_THAT(flatten.Run(module.get()), IsOkAndHolds(true));
HloConstantFolding constant_folding;
EXPECT_THAT(constant_folding.Run(module.get()), IsOkAndHolds(true));
EXPECT_THAT(RunFileCheck(module->ToString(), hlo_string),
IsOkAndHolds(true));
}
}
TEST_F(HloConstantFoldingTest, ComputationCalledThriceOnTwoDifferentConstants) {
const std::string hlo_string = R"(
HloModule test
// CHECK-LABEL: %subcomp (
// CHECK: constant({15, 15})
// CHECK-LABEL: %subcomp.clone (
// CHECK: constant({25, 25})
// CHECK-LABEL: %subcomp.clone.1 (
// CHECK: constant({25, 25})
subcomp {
p0 = s32[2] parameter(0)
p1 = s32[2] parameter(1)
c = s32[2] constant({5, 5})
add0 = s32[2] add(p0, c)
ROOT add1 = s32[2] add(p1, add0)
}
// CHECK-LABEL: ENTRY %entry (
// CHECK: %call1 = s32[2]{0} call(%c10, %p0.1), to_apply=%subcomp
// CHECK: %call2 = s32[2]{0} call(%c20, %p0.1), to_apply=%subcomp.clone
// CHECK: %call3 = s32[2]{0} call(%c20, %p0.1), to_apply=%subcomp.clone.1
ENTRY entry {
p0.1 = s32[2] parameter(0)
c10 = s32[2] constant({10, 10})
c20 = s32[2] constant({20, 20})
call1 = s32[2] call(c10, p0.1), to_apply=subcomp
call2 = s32[2] call(c20, p0.1), to_apply=subcomp
call3 = s32[2] call(c20, p0.1), to_apply=subcomp
add1 = s32[2] add(call1, call2)
ROOT out = s32[2] add(add1, call3)
})";
{
ASSERT_OK_AND_ASSIGN(auto module, ParseAndReturnVerifiedModule(hlo_string));
HloConstantFolding constant_folding;
EXPECT_THAT(constant_folding.Run(module.get()), IsOkAndHolds(true));
}
{
ASSERT_OK_AND_ASSIGN(auto module, ParseAndReturnVerifiedModule(hlo_string));
FlattenCallGraph flatten;
EXPECT_THAT(flatten.Run(module.get()), IsOkAndHolds(true));
HloConstantFolding constant_folding;
EXPECT_THAT(constant_folding.Run(module.get()), IsOkAndHolds(true));
EXPECT_THAT(RunFileCheck(module->ToString(), hlo_string),
IsOkAndHolds(true));
}
}
TEST_F(HloConstantFoldingTest,
ComputationCalledTwiceOnDifferentConstantsSingleArgument) {
const std::string hlo_string = R"(
HloModule test
// CHECK-LABEL: %subcomp (
// CHECK: constant({5, 5})
subcomp {
p0 = s32[2] parameter(0)
c = s32[2] constant({5, 5})
add0 = s32[2] add(p0, c)
ROOT add1 = s32[2] add(add0, add0)
}
// CHECK-LABEL: ENTRY %entry (
// CHECK: ROOT {{.*}} = s32[2]{0} constant({80, 80})
ENTRY entry {
c10 = s32[2] constant({10, 10})
c20 = s32[2] constant({20, 20})
call1 = s32[2] call(c10), to_apply=subcomp
call2 = s32[2] call(c20), to_apply=subcomp
ROOT out = s32[2] add(call1, call2)
})";
ASSERT_OK_AND_ASSIGN(auto module, ParseAndReturnVerifiedModule(hlo_string));
HloConstantFolding constant_folding;
EXPECT_THAT(constant_folding.Run(module.get()), IsOkAndHolds(true));
EXPECT_THAT(RunFileCheck(module->ToString(), hlo_string), IsOkAndHolds(true));
}
TEST_F(HloConstantFoldingTest,
ComputationCalledTwiceOnConstantAndOnNonConstant) {
const std::string hlo_string = R"(
HloModule test
// CHECK-LABEL: %subcomp (
// CHECK: constant({15, 15})
// CHECK-LABEL: %subcomp.clone (
// CHECK: constant({5, 5})
subcomp {
p0 = s32[2] parameter(0)
p1 = s32[2] parameter(1)
c = s32[2] constant({5, 5})
add0 = s32[2] add(p0, c)
ROOT add1 = s32[2] add(p1, add0)
}
// CHECK-LABEL: ENTRY %entry (
// CHECK: %call1 = s32[2]{0} call(%c10, %p0.1), to_apply=%subcomp
// CHECK: %call2 = s32[2]{0} call(%p0.1, %p1.1), to_apply=%subcomp.clone
// CHECK: ROOT %out = s32[2]{0} add(%call1, %call2)
ENTRY entry {
p0.1 = s32[2] parameter(0)
p1.1 = s32[2] parameter(1)
c10 = s32[2] constant({10, 10})
call1 = s32[2] call(c10, p0.1), to_apply=subcomp
call2 = s32[2] call(p0.1, p1.1), to_apply=subcomp
ROOT out = s32[2] add(call1, call2)
})";
{
ASSERT_OK_AND_ASSIGN(auto module, ParseAndReturnVerifiedModule(hlo_string));
HloConstantFolding constant_folding;
EXPECT_THAT(constant_folding.Run(module.get()), IsOkAndHolds(true));
}
{
ASSERT_OK_AND_ASSIGN(auto module, ParseAndReturnVerifiedModule(hlo_string));
FlattenCallGraph flatten;
EXPECT_THAT(flatten.Run(module.get()), IsOkAndHolds(true));
HloConstantFolding constant_folding;
EXPECT_THAT(constant_folding.Run(module.get()), IsOkAndHolds(true));
EXPECT_THAT(RunFileCheck(module->ToString(), hlo_string),
IsOkAndHolds(true));
}
}
TEST_F(HloConstantFoldingTest,
ComputationCalledTwiceOnConstantAndOnNonConstantSingleArgument) {
const std::string hlo_string = R"(
HloModule test
// CHECK-LABEL: %subcomp (
// CHECK: constant({5, 5})
subcomp {
p0 = s32[2] parameter(0)
c = s32[2] constant({5, 5})
add0 = s32[2] add(p0, c)
ROOT add1 = s32[2] add(add0, add0)
}
// CHECK-LABEL: ENTRY %entry (
// CHECK: %constant = s32[2]{0} constant({30, 30})
// CHECK: %call2 = s32[2]{0} call(%p0.1), to_apply=%subcomp
// CHECK: ROOT %out = s32[2]{0} add(%constant, %call2)
ENTRY entry {
p0.1 = s32[2] parameter(0)
c10 = s32[2] constant({10, 10})
call1 = s32[2] call(c10), to_apply=subcomp
call2 = s32[2] call(p0.1), to_apply=subcomp
ROOT out = s32[2] add(call1, call2)
})";
ASSERT_OK_AND_ASSIGN(auto module, ParseAndReturnVerifiedModule(hlo_string));
HloConstantFolding constant_folding;
EXPECT_THAT(constant_folding.Run(module.get()), IsOkAndHolds(true));
EXPECT_THAT(RunFileCheck(module->ToString(), hlo_string), IsOkAndHolds(true));
}
TEST_F(HloConstantFoldingTest, ComputationCalledTwiceOnSameConstants) {
const std::string hlo_string = R"(
HloModule test
// CHECK-LABEL: %subcomp (
// CHECK: constant({15, 15})
subcomp {
p0 = s32[2] parameter(0)
p1 = s32[2] parameter(1)
c = s32[2] constant({5, 5})
add0 = s32[2] add(p0, c)
ROOT add1 = s32[2] add(p1, add0)
}
// CHECK-LABEL: ENTRY %entry (
// CHECK: %call1 = s32[2]{0} call(%c10, %p0.1), to_apply=%subcomp
// CHECK: %call2 = s32[2]{0} call(%c10, %p1.1), to_apply=%subcomp
// CHECK: ROOT %out = s32[2]{0} add(%call1, %call2)
ENTRY entry {
p0.1 = s32[2] parameter(0)
p1.1 = s32[2] parameter(1)
c10 = s32[2] constant({10, 10})
call1 = s32[2] call(c10, p0.1), to_apply=subcomp
call2 = s32[2] call(c10, p1.1), to_apply=subcomp
ROOT out = s32[2] add(call1, call2)
})";
ASSERT_OK_AND_ASSIGN(auto module, ParseAndReturnVerifiedModule(hlo_string));
HloConstantFolding constant_folding;
EXPECT_THAT(constant_folding.Run(module.get()), IsOkAndHolds(true));
EXPECT_THAT(RunFileCheck(module->ToString(), hlo_string), IsOkAndHolds(true));
}
TEST_F(HloConstantFoldingTest,
ComputationCalledTwiceOnSameConstantsSingleArgument) {
const std::string hlo_string = R"(
HloModule test
// CHECK-LABEL: %subcomp (
// CHECK: constant({5, 5})
subcomp {
p0 = s32[2] parameter(0)
c = s32[2] constant({5, 5})
add0 = s32[2] add(p0, c)
ROOT add1 = s32[2] add(add0, add0)
}
// CHECK-LABEL: ENTRY %entry (
// CHECK: ROOT {{.*}} = s32[2]{0} constant({60, 60})
ENTRY entry {
p0.1 = s32[2] parameter(0)
c10 = s32[2] constant({10, 10})
call1 = s32[2] call(c10), to_apply=subcomp
call2 = s32[2] call(c10), to_apply=subcomp
ROOT out = s32[2] add(call1, call2)
})";
ASSERT_OK_AND_ASSIGN(auto module, ParseAndReturnVerifiedModule(hlo_string));
HloConstantFolding constant_folding;
EXPECT_THAT(constant_folding.Run(module.get()), IsOkAndHolds(true));
EXPECT_THAT(RunFileCheck(module->ToString(), hlo_string), IsOkAndHolds(true));
}
TEST_F(HloConstantFoldingTest, ComputationCalledOnceOnConstants) {
const std::string hlo_string = R"(
HloModule test
// CHECK-LABEL: %subcomp (
// CHECK: constant({15, 15})
subcomp {
p0 = s32[2] parameter(0)
p1 = s32[2] parameter(1)
c = s32[2] constant({5, 5})
add0 = s32[2] add(p0, c)
ROOT add1 = s32[2] add(p1, add0)
}
// CHECK-LABEL: ENTRY %entry (
// CHECK: %call1 = s32[2]{0} call(%c10, %p0.1), to_apply=%subcomp
// CHECK: ROOT %out = s32[2]{0} negate(%call1)
ENTRY entry {
p0.1 = s32[2] parameter(0)
p1.1 = s32[2] parameter(1)
c10 = s32[2] constant({10, 10})
call1 = s32[2] call(c10, p0.1), to_apply=subcomp
ROOT out = s32[2] negate(call1)
})";
ASSERT_OK_AND_ASSIGN(auto module, ParseAndReturnVerifiedModule(hlo_string));
HloConstantFolding constant_folding;
EXPECT_THAT(constant_folding.Run(module.get()), IsOkAndHolds(true));
EXPECT_THAT(RunFileCheck(module->ToString(), hlo_string), IsOkAndHolds(true));
}
TEST_F(HloConstantFoldingTest, NestedCallOnConstant) {
const std::string hlo_string = R"(
HloModule test
// CHECK-LABEL: %bar (
// CHECK-NEXT: %p0.bar = s32[2]{0} parameter(0)
// CHECK-NEXT: %constant.1 = s32[2]{0} constant({10, 10})
// CHECK-NEXT: %p1.bar = s32[2]{0} parameter(1)
// CHECK-NEXT: ROOT %out.bar = s32[2]{0} add(%constant.1, %p1.bar)
bar {
p0.bar = s32[2] parameter(0)
p1.bar = s32[2] parameter(1)
negate.bar = s32[2] negate(p0.bar)
ROOT out.bar = s32[2] add(negate.bar, p1.bar)
}
// CHECK-LABEL: %foo (
// CHECK-NEXT: %p0.foo = s32[2]{0} parameter(0)
// CHECK-NEXT: %constant = s32[2]{0} constant({-10, -10})
// CHECK-NEXT: %p1.foo = s32[2]{0} parameter(1)
// CHECK-NEXT: %call.bar = s32[2]{0} call(%constant, %p1.foo), to_apply=%bar
// CHECK-NEXT: ROOT %out.foo = s32[2]{0} negate(%call.bar)
foo {
p0.foo = s32[2] parameter(0)
p1.foo = s32[2] parameter(1)
negate.foo = s32[2] negate(p0.foo)
call.bar = s32[2] call(negate.foo, p1.foo), to_apply=bar
ROOT out.foo = s32[2] negate(call.bar)
}
// CHECK-LABEL: ENTRY %entry (
// CHECK-NEXT: %c10 = s32[2]{0} constant({10, 10})
// CHECK-NEXT: %p0 = s32[2]{0} parameter(0)
// CHECK-NEXT: %call.foo = s32[2]{0} call(%c10, %p0), to_apply=%foo
// CHECK-NEXT: ROOT %out = s32[2]{0} negate(%call.foo)
ENTRY entry {
p0 = s32[2] parameter(0)
c10 = s32[2] constant({10, 10})
call.foo = s32[2] call(c10, p0), to_apply=foo
ROOT out = s32[2] negate(call.foo)
})";
ASSERT_OK_AND_ASSIGN(auto module, ParseAndReturnVerifiedModule(hlo_string));
HloConstantFolding constant_folding;
EXPECT_THAT(constant_folding.Run(module.get()), IsOkAndHolds(true));
EXPECT_THAT(RunFileCheck(module->ToString(), hlo_string), IsOkAndHolds(true));
}
TEST_F(HloConstantFoldingTest, NestedCallAndFooCalledTwice) {
const std::string hlo_string = R"(
HloModule test
// CHECK-LABEL: %bar (
// CHECK-NEXT: %p0.bar = s32[2]{0} parameter(0)
// CHECK-NEXT: %negate.bar = s32[2]{0} negate(%p0.bar)
// CHECK-NEXT: ROOT %out.bar = s32[2]{0} add(%negate.bar, %negate.bar)
bar {
p0.bar = s32[2] parameter(0)
negate.bar = s32[2] negate(p0.bar)
ROOT out.bar = s32[2] add(negate.bar, negate.bar)
}
// CHECK-LABEL: %foo (
// CHECK-NEXT: %p0.foo = s32[2]{0} parameter(0)
// CHECK-NEXT: %constant.2 = s32[2]{0} constant({20, 20})
// CHECK-NEXT: %p1.foo = s32[2]{0} parameter(1)
// CHECK-NEXT: ROOT %out.foo = s32[2]{0} add(%constant.2, %p1.foo)
foo {
p0.foo = s32[2] parameter(0)
p1.foo = s32[2] parameter(1)
negate.foo = s32[2] negate(p0.foo)
call.bar = s32[2] call(negate.foo), to_apply=bar
ROOT out.foo = s32[2] add(call.bar, p1.foo)
}
// CHECK-LABEL: ENTRY %entry (
// CHECK-NEXT: %c10 = s32[2]{0} constant({10, 10})
// CHECK-NEXT: %p0 = s32[2]{0} parameter(0)
// CHECK-NEXT: %call.0 = s32[2]{0} call(%c10, %p0), to_apply=%foo
// CHECK-NEXT: %constant = s32[2]{0} constant({30, 30})
// CHECK-NEXT: ROOT %out = s32[2]{0} add(%call.0, %constant)
ENTRY entry {
p0 = s32[2] parameter(0)
c10 = s32[2] constant({10, 10})
call.0 = s32[2] call(c10, p0), to_apply=foo
call.1 = s32[2] call(c10, c10), to_apply=foo
ROOT out = s32[2] add(call.0, call.1)
})";
ASSERT_OK_AND_ASSIGN(auto module, ParseAndReturnVerifiedModule(hlo_string));
HloConstantFolding constant_folding;
EXPECT_THAT(constant_folding.Run(module.get()), IsOkAndHolds(true));
EXPECT_THAT(RunFileCheck(module->ToString(), hlo_string), IsOkAndHolds(true));
}
TEST_F(HloConstantFoldingTest, NestedCallWhereOriginalCallerIsFolded) {
const std::string hlo_string = R"(
HloModule test
// CHECK-LABEL: %bar (
// CHECK: %constant.1 = s32[2]{0} constant({3, 3})
bar {
p0.bar = s32[2] parameter(0)
c1 = s32[2] constant({1, 1})
c2 = s32[2] constant({2, 2})
c3 = s32[2] add(c1, c2)
ROOT out.bar = s32[2] add(p0.bar, c3)
}
foo {
p0.foo = s32[2] parameter(0)
p1.foo = s32[2] parameter(1)
call.bar = s32[2] call(p0.foo), to_apply=bar
ROOT out.foo = s32[2] add(call.bar, p1.foo)
}
ENTRY entry {
p0 = s32[2] parameter(0)
p1 = s32[2] parameter(1)
c10 = s32[2] constant({10, 10})
call.0 = s32[2] call(c10, p0), to_apply=foo
call.1 = s32[2] call(p1, p0), to_apply=foo
ROOT out = s32[2] add(call.0, call.1)
})";
ASSERT_OK_AND_ASSIGN(auto module, ParseAndReturnVerifiedModule(hlo_string));
HloConstantFolding constant_folding;
EXPECT_THAT(constant_folding.Run(module.get()), IsOkAndHolds(true));
EXPECT_THAT(RunFileCheck(module->ToString(), hlo_string), IsOkAndHolds(true));
}
} // namespace
} // namespace xla