mirror of
https://github.com/tensorflow/tensorflow.git
synced 2026-09-28 13:23:37 +08:00
committed by
TensorFlower Gardener
parent
8c9cfc9f5a
commit
d9ecfe9451
+3
-1
@@ -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
|
||||
|
||||
+447
-7
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user