chore: import upstream snapshot with attribution

This commit is contained in:
wehub-resource-sync
2026-07-13 12:40:42 +08:00
commit e25996e7db
15472 changed files with 3536181 additions and 0 deletions
+8
View File
@@ -0,0 +1,8 @@
if(WITH_TESTING AND WITH_CINN)
paddle_test(map_expr_test SRCS map_expr_test.cc)
set_tests_properties(map_expr_test PROPERTIES LABELS "RUN_TYPE=CINN")
paddle_test(test_data_dependency_graph SRCS data_dependency_graph_test.cc)
paddle_test(test_index_expr SRCS index_expr_test.cc)
paddle_test(test_iter_simplify SRCS iter_simplify_test.cc)
paddle_test(merge_block_utils_test SRCS merge_block_utils_test.cc)
endif()
@@ -0,0 +1,124 @@
// Copyright (c) 2024 PaddlePaddle Authors. All Rights Reserved.
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
#include <glog/logging.h>
#include <gtest/gtest.h>
#include <sstream>
#include "paddle/cinn/ir/ir_analyzer/data_dependency_graph.h"
#include "paddle/cinn/ir/ir_printer.h"
#include "paddle/cinn/ir/stmt.h"
namespace cinn {
namespace ir {
namespace {
using ir::analyzer::DataDependencyGraph;
using ir::analyzer::DepKind;
class TestDataDependencyGraph : public ::testing::Test {
public:
void SetUp() override {
const std::vector<ir::Expr> shape = {};
tensor_a =
ir::_Tensor_::Make(common::UniqName("A"), common::Bool(), shape, shape);
tensor_b =
ir::_Tensor_::Make(common::UniqName("B"), common::Bool(), shape, shape);
tensor_c =
ir::_Tensor_::Make(common::UniqName("C"), common::Bool(), shape, shape);
tensor_d =
ir::_Tensor_::Make(common::UniqName("D"), common::Bool(), shape, shape);
tensor_a->WithBuffer("global", "_" + tensor_a->name + "_temp_buffer");
tensor_b->WithBuffer("global", "_" + tensor_b->name + "_temp_buffer");
tensor_c->WithBuffer("global", "_" + tensor_c->name + "_temp_buffer");
tensor_d->WithBuffer("global", "_" + tensor_d->name + "_temp_buffer");
var_x = ir::Var(ir::Expr(1), ir::Expr(INT32_MAX), "x");
load_a = ir::Load::Make(ir::Expr(tensor_a), {});
load_b = ir::Load::Make(ir::Expr(tensor_b), {});
load_c = ir::Load::Make(ir::Expr(tensor_c), {});
};
ir::Var var_x;
ir::Tensor tensor_a, tensor_b, tensor_c, tensor_d;
ir::Expr load_a, load_b, load_c;
};
TEST_F(TestDataDependencyGraph, TensorDep) {
// A[0] = B[0]
auto store_a_b = ir::stmt::Store(ir::Expr(tensor_a), load_b, {});
// B[0] = C[0]
auto store_b_c = ir::stmt::Store(ir::Expr(tensor_b), load_c, {});
// C[0] = A[0]
auto store_c_a = ir::stmt::Store(ir::Expr(tensor_c), load_a, {});
const std::vector<ir::stmt::StmtRef> stmts = {
store_a_b, store_b_c, store_c_a};
const auto &dep_graph = DataDependencyGraph(stmts);
dep_graph.Print();
EXPECT_EQ(dep_graph.HasDependency(stmts[0], stmts[1]), DepKind::DEP);
EXPECT_EQ(dep_graph.HasDependency(stmts[0], stmts[2]), DepKind::DEP);
EXPECT_EQ(dep_graph.HasDependency(stmts[1], stmts[2]), DepKind::DEP);
}
TEST_F(TestDataDependencyGraph, VarDep) {
// x = A[0]
auto let_x_a = ir::stmt::Let(ir::Expr(var_x), load_a);
// B[0] = x
auto store_b_x = ir::stmt::Store(ir::Expr(tensor_b), ir::Expr(var_x), {});
const std::vector<ir::stmt::StmtRef> stmts = {let_x_a, store_b_x};
const auto &dep_graph = DataDependencyGraph(stmts);
dep_graph.Print();
EXPECT_EQ(dep_graph.HasDependency(stmts[0], stmts[1]), DepKind::DEP);
}
TEST_F(TestDataDependencyGraph, TensorNoDep) {
// A[0] = B[0]
auto store_a_b = ir::stmt::Store(ir::Expr(tensor_a), load_b, {});
// C[0] = B[0]
auto store_c_b = ir::stmt::Store(ir::Expr(tensor_c), load_b, {});
// D[0] = B[0]
auto store_d_b = ir::stmt::Store(ir::Expr(tensor_d), load_b, {});
const std::vector<ir::stmt::StmtRef> stmts = {
store_a_b, store_c_b, store_d_b};
const auto &dep_graph = DataDependencyGraph(stmts);
dep_graph.Print();
EXPECT_EQ(dep_graph.HasDependency(stmts[0], stmts[1]), DepKind::NO_DEP);
EXPECT_EQ(dep_graph.HasDependency(stmts[0], stmts[2]), DepKind::NO_DEP);
EXPECT_EQ(dep_graph.HasDependency(stmts[1], stmts[2]), DepKind::NO_DEP);
}
TEST_F(TestDataDependencyGraph, VarNoDep) {
// x = A[0]
auto let_x_a = ir::stmt::Let(ir::Expr(var_x), load_a);
// B[0] = x
auto store_b_x = ir::stmt::Store(ir::Expr(tensor_b), ir::Expr(var_x), {});
// C[0] = x
auto store_c_x = ir::stmt::Store(ir::Expr(tensor_c), ir::Expr(var_x), {});
// D[0] = x
auto store_d_x = ir::stmt::Store(ir::Expr(tensor_d), ir::Expr(var_x), {});
const std::vector<ir::stmt::StmtRef> stmts = {
let_x_a, store_b_x, store_c_x, store_d_x};
const auto &dep_graph = DataDependencyGraph(stmts);
dep_graph.Print();
EXPECT_EQ(dep_graph.HasDependency(stmts[1], stmts[2]), DepKind::NO_DEP);
EXPECT_EQ(dep_graph.HasDependency(stmts[1], stmts[3]), DepKind::NO_DEP);
EXPECT_EQ(dep_graph.HasDependency(stmts[2], stmts[3]), DepKind::NO_DEP);
}
} // namespace
} // namespace ir
} // namespace cinn
+688
View File
@@ -0,0 +1,688 @@
// Copyright (c) 2024 CINN Authors. All Rights Reserved.
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
#include <glog/logging.h>
#include <gtest/gtest.h>
#include "paddle/cinn/common/integer_set.h"
#include "paddle/cinn/common/simplify_special_pattern.h"
#include "paddle/cinn/ir/ir.h"
#include "paddle/cinn/ir/ir_base.h"
#include "paddle/cinn/ir/ir_mutator.h"
#include "paddle/cinn/ir/op/ir_operators.h"
#include "paddle/cinn/optim/simplify_util.h"
namespace cinn {
namespace common {
using optim::ChangeSeqOfDivMod;
using optim::CheckPattern;
using optim::ConstructIndexExprByNodeType;
using optim::MatchPattern;
using optim::ParseExpressionFromString;
class TestIndexExpr : public ::testing::Test {
public:
void SetUp() override {
S4 = ir::Var(ir::Expr(static_cast<int64_t>(1)), ir::Expr(INT32_MAX), "S4")
.set_index(true);
S5 = ir::Var(ir::Expr(static_cast<int64_t>(1)), ir::Expr(INT32_MAX), "S5")
.set_index(true);
S6 = ir::Var(ir::Expr(static_cast<int64_t>(1)), ir::Expr(INT32_MAX), "S6")
.set_index(true);
S7 = ir::Var(ir::Expr(static_cast<int64_t>(1)), ir::Expr(INT32_MAX), "S7")
.set_index(true);
S8 = ir::Var(ir::Expr(static_cast<int64_t>(1)), ir::Expr(INT32_MAX), "S8")
.set_index(true);
S9 = ir::Var(ir::Expr(static_cast<int64_t>(1)), ir::Expr(INT32_MAX), "S9")
.set_index(true);
f = ir::Var(ir::Expr(static_cast<int64_t>(1)), ir::Expr(INT32_MAX), "f");
};
ir::Var S4, S5, S6, S7, S8, S9, f;
};
TEST_F(TestIndexExpr, IndexExpr_0) {
ir::IndexExpr a(14);
ir::IndexExpr b(7);
Expr d(6);
ir::Expr c0 = a + b;
ir::Expr c1 = a - b;
ir::Expr c2 = a * b;
ir::Expr c3 = a / b;
ir::Expr c4 = a % b;
ir::Expr c5 = a / d.as_index();
ir::Expr c6 = a % d.as_index();
EXPECT_EQ(c0, Expr(21));
EXPECT_EQ(c1, Expr(7));
EXPECT_EQ(c2, Expr(98));
EXPECT_EQ(c3, Expr(2));
EXPECT_EQ(c4, Expr(0));
EXPECT_EQ(c5, Expr(2));
EXPECT_EQ(c6, Expr(2));
}
TEST_F(TestIndexExpr, IndexExpr_1) {
auto test = S6 * S7;
ir::IndexExpr e1 = (S5 * ((S4 * (S5 * (S6 * S7))) / S5));
ir::IndexExpr e2 = (S4 * (S5 * (S6 * S7))) / S5;
ir::IndexExpr e3 = (S4 * S5) / S5;
ir::IndexExpr e4 = (S4 * (S5 * (S6 * S7)) + S5) / S5;
ir::IndexExpr e5 = (S4 * (S5 * (S6 * S7)) + 2 * S5) / S5;
ir::IndexExpr e6 = (S4 * (S5 * (S6 * S7)) + S5 / S6) / S5;
ir::IndexExpr e7 = (S4 * (S5 * (S6 * S7)) + 2 * S5 / S6) / S5;
EXPECT_EQ(e1.Normalize(), ir::IndexExpr((S6 * S7) * S4 * S5));
EXPECT_EQ(e2.Normalize(), ir::IndexExpr((S6 * S7) * S4));
EXPECT_EQ(e3.Normalize(), ir::IndexExpr(S4));
EXPECT_EQ(e4.Normalize(), ir::IndexExpr(((S6 * S7) * S4) + 1));
EXPECT_EQ(e5.Normalize(), ir::IndexExpr(((S6 * S7) * S4) + 2));
EXPECT_EQ(e6.Normalize(), ir::IndexExpr(((S6 * S7) * S4) + (1 / S6)));
EXPECT_EQ(e7.Normalize(), ir::IndexExpr(((S6 * S7) * S4) + (2 / S6)));
}
TEST_F(TestIndexExpr, IndexExpr_2) {
ir::Expr q1 = S4;
ir::Expr q2 = S4;
ir::Expr q3 = S4 + S5;
ir::Expr q4 = S5 + S4;
ir::Expr q5 = S4 * 2 + S5 / 4;
ir::Expr q6 = S5 / 4 + S4 * 2;
ir::Expr q7 = S4 + S5 + S6;
ir::Expr q8 = S5 + (S4 + S6);
ir::Expr q9 = S4 + (S5 + S7 / 4 + S6 * 2);
ir::Expr q10 = S5 + (S4 + S6 * 2 + S7 / 4);
ir::Expr q11 = (S7 + S5) + (S4 + S6);
ir::Expr q12 = (S4 + S5) + (S6 + S7);
ir::Expr q13 = (S4 + S5) * 3 + (S6 / 2 + S7) * 2;
ir::Expr q14 = (S6 / 2 + S7) * 2 + (S4 + S5) * 3;
ir::Expr q15 = (S4 + S5 * 2) * 3 + (S6 / 2 + S7) * 2;
ir::Expr q16 = (S6 / 2 + S7) * 2 + (S4 + S5 * 2) * 3;
ir::Expr q17 = (S4 + S5 * 2) * 3 + (S6 / 2 + S7) * 2 + S4;
ir::Expr q18 = (S6 / 2 + S7) * 2 + (S4 + S5 * 2) * 3 + S4;
ir::Expr q19 = (S4 + S5 * 2) * 3 + (S6 / 2 + S7) * 2 + S4;
ir::Expr q20 = (S6 / 2 + S7) * 2 + (S4 + S5 * 2) * 3 + S5;
EXPECT_EQ(q1.as_index().Normalize(), q2.as_index().Normalize());
EXPECT_EQ(q3.as_index().Normalize(), q4.as_index().Normalize());
EXPECT_EQ(q5.as_index().Normalize(), q6.as_index().Normalize());
EXPECT_EQ(q7.as_index().Normalize(), q8.as_index().Normalize());
EXPECT_EQ(q9.as_index().Normalize(), q10.as_index().Normalize());
EXPECT_EQ(q11.as_index().Normalize(), q12.as_index().Normalize());
EXPECT_EQ(q13.as_index().Normalize(), q14.as_index().Normalize());
EXPECT_EQ(q15.as_index().Normalize(), q16.as_index().Normalize());
EXPECT_EQ(q17.as_index().Normalize(), q18.as_index().Normalize());
EXPECT_NE(q19.as_index().Normalize(), q20.as_index().Normalize());
}
TEST_F(TestIndexExpr, IndexExpr_3) {
// `Add` corner cases
ir::Expr q1 = S4 / S5 * S5 + S4 % S5;
ir::Expr q2 = (S4 + S5) / S6 * S6 + (S4 + S5) % S6;
ir::Expr q3 = S4 / (S5 + S6) * (S5 + S6) + S4 % (S5 + S6);
ir::Expr q4 = (S4 + S5) / (S6 + S7) * (S6 + S7) + (S4 + S5) % (S6 + S7);
ir::Expr q5 = (S4 + S5) / 5 * 5 + (S4 + S5) * 11 % 5;
ir::Expr q14 = (S4 + S5) / (S6 * S7) * S6 * S7 + (S4 + S5) % (S6 * S7);
ir::Expr q15 =
(S4 * 256 + S5 + S6 * 1024) % 25088 / 512 * 512 + (S4 * 256 + S5) % 512;
ir::Expr q16 =
((S4 * 256 + S5) / S6 / S7 * S7 + (S4 * 256 + S5) / S6 % S7) * S6 +
(S4 * 256 + S5) % S6;
ir::Expr q17 = S4 / (S5 * S6) * S6 + S4 % (S5 * S6) / S5;
ir::Expr q18 = (S4 * 1024 + S5 * 256 + S6) / 2097152 * 32 +
(S4 * 1024 + S5 * 256 + S6) % 2097152 / 65536;
// `Div` corner cases
ir::Expr q6 = (S4 % S5 - S4) / S5;
ir::Expr q7 = (S4 - S4 % S5) / S5;
ir::Expr q8 = ((S4 + S5) % S6 - S4 - S5) / S6;
ir::Expr q9 = (S4 + S5 - (S4 + S5) % S6) / S6;
// `Mod` corner cases
ir::Expr q10 = (S4 % S5 - S4) % S5;
ir::Expr q11 = (S4 - S4 % S5) % S5;
ir::Expr q12 = ((S4 + S5) % S6 - S4 - S5) % S6;
ir::Expr q13 = (S4 + S5 - (S4 + S5) % S6) % S6;
EXPECT_EQ(q1.as_index().Normalize(), ir::IndexExpr(S4));
EXPECT_EQ(q2.as_index().Normalize(), ir::IndexExpr(S4 + S5));
EXPECT_EQ(q3.as_index().Normalize(), ir::IndexExpr(S4));
EXPECT_EQ(q4.as_index().Normalize(), ir::IndexExpr(S4 + S5));
EXPECT_EQ(q5.as_index().Normalize(), ir::IndexExpr(S4 + S5));
EXPECT_EQ(q6.as_index().Normalize(), ir::IndexExpr((S4 / S5) * (-1)));
EXPECT_EQ(q7.as_index().Normalize(), ir::IndexExpr(S4 / S5));
EXPECT_EQ(q8.as_index().Normalize(), ir::IndexExpr(((S4 + S5) / S6) * (-1)));
EXPECT_EQ(q9.as_index().Normalize(), ir::IndexExpr((S4 + S5) / S6));
EXPECT_EQ(q10.as_index().Normalize(), ir::IndexExpr(0));
EXPECT_EQ(q11.as_index().Normalize(), ir::IndexExpr(0));
EXPECT_EQ(q12.as_index().Normalize(), ir::IndexExpr(0));
EXPECT_EQ(q13.as_index().Normalize(), ir::IndexExpr(0));
EXPECT_EQ(q14.as_index().Normalize(), ir::IndexExpr(S4 + S5));
EXPECT_EQ(q15.as_index().Normalize(),
ir::IndexExpr((S4 * 256 + S5 + S6 * 1024)) % 25088);
EXPECT_EQ(q16.as_index().Normalize(ir::IndexExpr::OptLevel::kLevel2),
ir::IndexExpr(S4 * 256 + S5));
EXPECT_EQ(q17.as_index().Normalize(), ir::IndexExpr(S4 / S5));
EXPECT_EQ(q18.as_index().Normalize(),
ir::IndexExpr((S4 * 1024 + S5 * 256 + S6) / 65536));
}
TEST_F(TestIndexExpr, Change_Seq_Of_Div_Mod) {
ir::Expr q1 = S4 / S5;
ir::Expr q2 = S4 % S5;
ir::Expr q3 = S4 / S5 % S6;
ir::Expr q4 = S4 / S5 % S6;
EXPECT_EQ(ChangeSeqOfDivMod(q1.as_index()), q1);
EXPECT_EQ(ChangeSeqOfDivMod(q2.as_index()), q2);
EXPECT_EQ(ChangeSeqOfDivMod(q3.as_index()), S4 % (S5 * S6) / S5);
}
TEST_F(TestIndexExpr, Test_ConstructIndexExprByNodeType) {
ir::Expr result_add = ConstructIndexExprByNodeType(
ir::IrNodeTy::Add, S4.as_index(), S5.as_index(), true);
ir::Expr result_sub = ConstructIndexExprByNodeType(
ir::IrNodeTy::Sub, S4.as_index(), S5.as_index(), false);
ir::Expr result_mul = ConstructIndexExprByNodeType(
ir::IrNodeTy::Mul, S4.as_index(), S5.as_index(), true);
ir::Expr result_div = ConstructIndexExprByNodeType(
ir::IrNodeTy::Div, S4.as_index(), S5.as_index(), true);
ir::Expr result_mod = ConstructIndexExprByNodeType(
ir::IrNodeTy::Mod, S4.as_index(), S5.as_index(), true);
ir::Expr result_min = ConstructIndexExprByNodeType(
ir::IrNodeTy::Min, S4.as_index(), S5.as_index(), false);
ir::Expr result_max = ConstructIndexExprByNodeType(
ir::IrNodeTy::Max, S4.as_index(), S5.as_index(), false);
EXPECT_EQ(result_add, S4 + S5);
EXPECT_EQ(result_sub, S4 - S5);
EXPECT_EQ(result_mul, S4 * S5);
EXPECT_EQ(result_div, S4 / S5);
EXPECT_EQ(result_mod, S4 % S5);
EXPECT_EQ(result_min, ir::Min::Make(S4, S5));
EXPECT_EQ(result_max, ir::Max::Make(S4, S5));
}
TEST_F(TestIndexExpr, Test_dynamic) {
ir::Expr q =
((((((((((((((((S7 * 1024) + S8) + (S9 * 4096)) / S6) % S5) * S6) +
((((S7 * 1024) + S8) + (S9 * 4096)) % S6)) +
(((((((((S7 * 1024) + S8) + (S9 * 4096)) / S6) / S5) % 640) %
S4) *
S6) *
S5)) +
(((((((((S7 * 1024) + S8) + (S9 * 4096)) / S6) / S5) / 640) *
S5) *
S6) *
S4)) /
((S5 * S6) * S4)) *
S4) +
(((((((((((S7 * 1024) + S8) + (S9 * 4096)) / S6) % S5) * S6) +
((((S7 * 1024) + S8) + (S9 * 4096)) % S6)) +
(((((((((S7 * 1024) + S8) + (S9 * 4096)) / S6) / S5) % 640) %
S4) *
S6) *
S5)) +
(((((((((S7 * 1024) + S8) + (S9 * 4096)) / S6) / S5) / 640) *
S5) *
S6) *
S4)) /
(S5 * S6)) %
S4)) *
S5) +
(((((((((((S7 * 1024) + S8) + (S9 * 4096)) / S6) % S5) * S6) +
((((S7 * 1024) + S8) + (S9 * 4096)) % S6)) +
(((((((((S7 * 1024) + S8) + (S9 * 4096)) / S6) / S5) % 640) % S4) *
S6) *
S5)) +
(((((((((S7 * 1024) + S8) + (S9 * 4096)) / S6) / S5) / 640) * S5) *
S6) *
S4)) /
S6) %
S5)) *
S6) +
((((((((((S7 * 1024) + S8) + (S9 * 4096)) / S6) % S5) * S6) +
((((S7 * 1024) + S8) + (S9 * 4096)) % S6)) +
(((((((((S7 * 1024) + S8) + (S9 * 4096)) / S6) / S5) % 640) % S4) *
S6) *
S5)) +
(((((((((S7 * 1024) + S8) + (S9 * 4096)) / S6) / S5) / 640) * S5) *
S6) *
S4)) %
S6));
ir::Expr q1 =
((((((f % ((S5 * S6) * 640)) % ((S5 * S6) * S4)) / (S5 * S6)) * S6) *
S5) +
(f % (S5 * S6)));
ir::Expr q2 = ((f % ((S5 * S6) * 640)) % ((S5 * S6) * S4)) % (S5 * S6);
ir::Expr q3 = (S5 * S6) * S4 / (S5 * S6);
ir::Expr q4 = (S5 * S6) * S4 % (S5 * S6);
ir::Expr q5 =
(((((((((((((f % ((S5 * S6) * 640)) % ((S5 * S6) * S4)) / (S5 * S6)) +
((f / ((S5 * S6) * 640)) * S4)) *
S5) *
S6) +
(f % (S5 * S6))) %
((S5 * S6) * S4)) /
(S5 * S6)) +
(((((((((f % ((S5 * S6) * 640)) % ((S5 * S6) * S4)) / (S5 * S6)) +
((f / ((S5 * S6) * 640)) * S4)) *
S5) *
S6) +
(f % (S5 * S6))) /
((S5 * S6) * S4)) *
S4)) *
S5) *
S6) +
(f % (S5 * S6)));
EXPECT_EQ(
q.as_index().Normalize(ir::IndexExpr::OptLevel::kLevel2),
((((((((S7 * 1024) + S8) + (S9 * 4096)) / ((S5 * S6) * 640)) * S5) * S6) *
S4) +
(((((S7 * 1024) + S8) + (S9 * 4096)) % ((S5 * S6) * 640)) %
((S5 * S6) * S4))));
EXPECT_EQ(q1.as_index().Normalize(ir::IndexExpr::OptLevel::kLevel2),
((f % ((S5 * S6) * 640)) % ((S5 * S6) * S4)));
EXPECT_EQ(q2.as_index().Normalize(ir::IndexExpr::OptLevel::kLevel2),
f % (S5 * S6));
EXPECT_EQ(q3.as_index().Normalize(ir::IndexExpr::OptLevel::kLevel2),
Expr(S4));
EXPECT_EQ(q4.as_index().Normalize(ir::IndexExpr::OptLevel::kLevel2), Expr(0));
EXPECT_EQ(q5.as_index().Normalize(ir::IndexExpr::OptLevel::kLevel2),
(((((f / ((S5 * S6) * 640)) * S4) * S5) * S6) +
((f % ((S5 * S6) * 640)) % ((S5 * S6) * S4))));
}
TEST_F(TestIndexExpr, CommonFactor) {
ir::Var S0 = ir::Var("S0");
ir::Var S1 = ir::Var("S1");
ir::Var S2 = ir::Var("S2");
ir::Var S3 = ir::Var("S3");
ir::Var S4 = ir::Var("S4");
ir::Var S5 = ir::Var("S5");
ir::Var S6 = ir::Var("S6");
ir::Var S7 = ir::Var("S7");
ir::Var S8 = ir::Var("S8");
ir::Var S9 = ir::Var("S9");
ir::Var S13 = ir::Var("S13");
ir::Var S17 = ir::Var("S17");
ir::Var S21 = ir::Var("S21");
ir::Var tx = ir::Var("tx");
ir::Var bx = ir::Var("bx");
ir::Expr q = ((((((((S1 + S13) + S17) + S21) + S5) + S9)) * S2) * S3);
ir::Expr q1 = (((((((S3 * S5) * S2) + ((S3 * S9) * S2)) + ((S3 * S21) * S2)) +
((S2 * S3) * S17)) +
((S2 * S3) * S13)) +
((S2 * S3) * S1));
ir::Expr q2 =
(((((((((f * 1024) + tx) + (bx * 4096)) %
((((((((((((((((((((((((((S3 * S5) * S2) * S0) +
(((S3 * S9) * S2) * S0)) +
(((S3 * S21) * S2) * S0)) +
(((S2 * S3) * S17) * S0)) +
(((S2 * S3) * S13) * S0)) +
(((S2 * S3) * S1) * S0)) /
4096) *
4096) +
((S3 * S5) * S2)) +
((S3 * S9) * S2)) +
((S3 * S21) * S2)) +
((S2 * S3) * S17)) +
((S2 * S3) * S13)) +
((S2 * S3) * S1)) +
4095) /
(((((((S3 * S5) * S2) + ((S3 * S9) * S2)) +
((S3 * S21) * S2)) +
((S2 * S3) * S17)) +
((S2 * S3) * S13)) +
((S2 * S3) * S1))) *
S3) *
S5) *
S2) +
(((((((((((((((((((((S3 * S5) * S2) * S0) +
(((S3 * S9) * S2) * S0)) +
(((S3 * S21) * S2) * S0)) +
(((S2 * S3) * S17) * S0)) +
(((S2 * S3) * S13) * S0)) +
(((S2 * S3) * S1) * S0)) /
4096) *
4096) +
((S3 * S5) * S2)) +
((S3 * S9) * S2)) +
((S3 * S21) * S2)) +
((S2 * S3) * S17)) +
((S2 * S3) * S13)) +
((S2 * S3) * S1)) +
4095) /
(((((((S3 * S5) * S2) + ((S3 * S9) * S2)) +
((S3 * S21) * S2)) +
((S2 * S3) * S17)) +
((S2 * S3) * S13)) +
((S2 * S3) * S1))) *
S3) *
S9) *
S2)) +
(((((((((((((((((((((S3 * S5) * S2) * S0) +
(((S3 * S9) * S2) * S0)) +
(((S3 * S21) * S2) * S0)) +
(((S2 * S3) * S17) * S0)) +
(((S2 * S3) * S13) * S0)) +
(((S2 * S3) * S1) * S0)) /
4096) *
4096) +
((S3 * S5) * S2)) +
((S3 * S9) * S2)) +
((S3 * S21) * S2)) +
((S2 * S3) * S17)) +
((S2 * S3) * S13)) +
((S2 * S3) * S1)) +
4095) /
(((((((S3 * S5) * S2) + ((S3 * S9) * S2)) +
((S3 * S21) * S2)) +
((S2 * S3) * S17)) +
((S2 * S3) * S13)) +
((S2 * S3) * S1))) *
S3) *
S21) *
S2)) +
(((((((((((((((((((((S3 * S5) * S2) * S0) +
(((S3 * S9) * S2) * S0)) +
(((S3 * S21) * S2) * S0)) +
(((S2 * S3) * S17) * S0)) +
(((S2 * S3) * S13) * S0)) +
(((S2 * S3) * S1) * S0)) /
4096) *
4096) +
((S3 * S5) * S2)) +
((S3 * S9) * S2)) +
((S3 * S21) * S2)) +
((S2 * S3) * S17)) +
((S2 * S3) * S13)) +
((S2 * S3) * S1)) +
4095) /
(((((((S3 * S5) * S2) + ((S3 * S9) * S2)) +
((S3 * S21) * S2)) +
((S2 * S3) * S17)) +
((S2 * S3) * S13)) +
((S2 * S3) * S1))) *
S2) *
S3) *
S17)) +
(((((((((((((((((((((S3 * S5) * S2) * S0) +
(((S3 * S9) * S2) * S0)) +
(((S3 * S21) * S2) * S0)) +
(((S2 * S3) * S17) * S0)) +
(((S2 * S3) * S13) * S0)) +
(((S2 * S3) * S1) * S0)) /
4096) *
4096) +
((S3 * S5) * S2)) +
((S3 * S9) * S2)) +
((S3 * S21) * S2)) +
((S2 * S3) * S17)) +
((S2 * S3) * S13)) +
((S2 * S3) * S1)) +
4095) /
(((((((S3 * S5) * S2) + ((S3 * S9) * S2)) +
((S3 * S21) * S2)) +
((S2 * S3) * S17)) +
((S2 * S3) * S13)) +
((S2 * S3) * S1))) *
S2) *
S3) *
S13)) +
(((((((((((((((((((((S3 * S5) * S2) * S0) +
(((S3 * S9) * S2) * S0)) +
(((S3 * S21) * S2) * S0)) +
(((S2 * S3) * S17) * S0)) +
(((S2 * S3) * S13) * S0)) +
(((S2 * S3) * S1) * S0)) /
4096) *
4096) +
((S3 * S5) * S2)) +
((S3 * S9) * S2)) +
((S3 * S21) * S2)) +
((S2 * S3) * S17)) +
((S2 * S3) * S13)) +
((S2 * S3) * S1)) +
4095) /
(((((((S3 * S5) * S2) + ((S3 * S9) * S2)) +
((S3 * S21) * S2)) +
((S2 * S3) * S17)) +
((S2 * S3) * S13)) +
((S2 * S3) * S1))) *
S2) *
S3) *
S1))) /
(((((((S3 * S5) * S2) + ((S3 * S9) * S2)) + ((S3 * S21) * S2)) +
((S2 * S3) * S17)) +
((S2 * S3) * S13)) +
((S2 * S3) * S1))) *
(((((S1 + S13) + S17) + S21) + S5) + S9)) *
S2) *
S3) +
((((f * 1024) + tx) + (bx * 4096)) %
(((((((S3 * S5) * S2) + ((S3 * S9) * S2)) + ((S3 * S21) * S2)) +
((S2 * S3) * S17)) +
((S2 * S3) * S13)) +
((S2 * S3) * S1))));
EXPECT_EQ(q.as_index().Normalize(ir::IndexExpr::OptLevel::kLevel2),
(((((((S1 + S13) + S17) + S21) + S5) + S9) * S2) * S3));
EXPECT_EQ(q1.as_index().Normalize(ir::IndexExpr::OptLevel::kLevel2),
(((((((S5 + S9) + S21) + S17) + S13) + S1) * S2) * S3));
EXPECT_EQ(
q2.as_index().Normalize(ir::IndexExpr::OptLevel::kLevel2),
((((f * 1024) + tx) + (bx * 4096)) %
((((((((((((((((S5 + S9) + S21) + S17) + S13) + S1) * S2) * S3) * S0) /
4096) *
4096) +
(((((((S5 + S9) + S21) + S17) + S13) + S1) * S2) * S3)) +
4095) /
(((((((S5 + S9) + S21) + S17) + S13) + S1) * S2) * S3)) *
S3) *
S2) *
(((((S5 + S9) + S21) + S17) + S13) + S1))));
}
TEST_F(TestIndexExpr, TestCheckPattern) {
ir::Var a = ir::Var("a");
ir::Var b = ir::Var("b");
ir::Var f = ir::Var("f");
ir::Var S0 = ir::Var("S0");
ir::Var S1 = ir::Var("S1");
ir::Var S2 = ir::Var("S2");
ir::Var S3 = ir::Var("S3");
ir::Var S4 = ir::Var("S4");
ir::Var S5 = ir::Var("S5");
ir::Var S6 = ir::Var("S6");
ir::Var S7 = ir::Var("S7");
ir::Var S8 = ir::Var("S8");
ir::Var S9 = ir::Var("S9");
ir::IndexExpr pattern = f / (a * b) * b + f % (a * b) / a;
ir::IndexExpr pattern1 = f / (a * b) * a + f % (a * b) / b;
ir::IndexExpr e = (S0 * (S1 + S2) + S1 * S2 + S2) / (S4 * S5) * S5 +
(S0 * (S1 + S2) + S1 * S2 + S2) % (S4 * S5) / S4;
ir::IndexExpr e1 = (S0 * (S1 + S2) + S1 * S2 + S2) / (S4 * S5) * S4 +
(S0 * (S1 + S2) + S1 * S2 + S2) % (S4 * S5) / S5;
std::unordered_map<std::string, ir::IndexExpr> map;
EXPECT_TRUE(CheckPattern(e, pattern, &map));
map.clear();
EXPECT_FALSE(CheckPattern(e, pattern1, &map));
map.clear();
EXPECT_FALSE(CheckPattern(e1, pattern, &map));
map.clear();
EXPECT_TRUE(CheckPattern(e1, pattern1, &map));
}
TEST_F(TestIndexExpr, ParseExpression) {
ir::Var a = ir::Var("a");
ir::Var b = ir::Var("b");
ir::Var a1 = ir::Var("a_1");
ir::Var b2 = ir::Var("b2");
ir::Expr e1 = a + b;
ir::Expr e2 = a - b;
ir::Expr e3 = a * b;
ir::Expr e4 = a / b;
ir::Expr e5 = a % b;
ir::Expr e6 = a + ir::Expr(20);
ir::Expr e7 = a - ir::Expr(10);
ir::Expr e8 = ir::Expr(5) * b;
ir::Expr e9 = ir::Expr(20) / b;
ir::Expr e10 = a % ir::Expr(3) + b;
ir::Expr e11 = (a + b) * (a - b);
ir::Expr e12 = (a + (b * a)) - (b / a);
ir::Expr e13 = (a + b) * (a - b) + (a / b) - (b % a);
ir::Expr e14 = a1 + b2;
ir::Expr e15 = a + b;
EXPECT_EQ(e1, ParseExpressionFromString("a + b"));
EXPECT_EQ(e2, ParseExpressionFromString("a - b"));
EXPECT_EQ(e3, ParseExpressionFromString("a * b"));
EXPECT_EQ(e4, ParseExpressionFromString("a / b"));
EXPECT_EQ(e5, ParseExpressionFromString("a % b"));
EXPECT_EQ(e6, ParseExpressionFromString("a + 20"));
EXPECT_EQ(e7, ParseExpressionFromString("a - 10"));
EXPECT_EQ(e8, ParseExpressionFromString("5 * b"));
EXPECT_EQ(e9, ParseExpressionFromString("20 / b"));
EXPECT_EQ(e10, ParseExpressionFromString("a % 3 + b"));
EXPECT_EQ(e11, ParseExpressionFromString("(a + b) * (a - b)"));
EXPECT_EQ(e12, ParseExpressionFromString("(a + (b * a)) - (b / a)"));
EXPECT_EQ(e13,
ParseExpressionFromString("(a + b) * (a - b) + (a / b) - (b % a)"));
EXPECT_EQ(e14, ParseExpressionFromString("a_1 + b2"));
EXPECT_EQ(e15, ParseExpressionFromString(" a + b "));
EXPECT_ANY_THROW(ParseExpressionFromString("a + #"));
EXPECT_ANY_THROW(ParseExpressionFromString("(a + b"));
EXPECT_ANY_THROW(ParseExpressionFromString(""));
}
TEST_F(TestIndexExpr, MatchPattern) {
ir::Var a = ir::Var("a");
ir::Var b = ir::Var("b");
ir::Var x = ir::Var("x");
ir::Var y = ir::Var("y");
ir::IndexExpr expr1 = a + b;
ir::IndexExpr expr2 = a * b;
ir::IndexExpr expr3 = a + (b * 10);
ir::IndexExpr expr4 = (a + b) * 10;
ir::IndexExpr expr5 = x + y;
ir::IndexExpr expr6 = x * y;
auto result1 = MatchPattern(expr1, "a + b", nullptr);
EXPECT_TRUE(result1.has_value());
EXPECT_EQ(result1->at("a"), a);
EXPECT_EQ(result1->at("b"), b);
auto result2 = MatchPattern(expr3, "a + (b * 10)", nullptr);
EXPECT_TRUE(result2.has_value());
EXPECT_EQ(result2->at("a"), a);
EXPECT_EQ(result2->at("b"), b);
auto result3 = MatchPattern(expr1, "a * b", nullptr);
EXPECT_FALSE(result3.has_value());
auto result4 = MatchPattern(expr3, "a + (b * 20)", nullptr);
EXPECT_FALSE(result4.has_value());
auto condition =
[](const std::unordered_map<std::string, ir::IndexExpr> &map) {
return map.at("a") == Expr(ir::Var("a")) &&
map.at("b") == Expr(ir::Var("b"));
};
auto result5 = MatchPattern(expr1, "a + b", condition);
EXPECT_TRUE(result5.has_value());
auto condition2 =
[](const std::unordered_map<std::string, ir::IndexExpr> &map) {
return map.at("a") == ir::Var("x") && map.at("b") == ir::Var("y");
};
auto result6 = MatchPattern(expr1, "a + b", condition2);
EXPECT_FALSE(result6.has_value());
auto result7 = MatchPattern(expr4, "(a + b) * 10", nullptr);
EXPECT_TRUE(result7.has_value());
EXPECT_EQ(result7->at("a"), a);
EXPECT_EQ(result7->at("b"), b);
auto result8 = MatchPattern(expr1, "x + y", nullptr);
EXPECT_TRUE(result8.has_value());
EXPECT_EQ(result8->at("x"), a);
EXPECT_EQ(result8->at("y"), b);
auto result9 = MatchPattern(expr6, "x * y", nullptr);
EXPECT_TRUE(result9.has_value());
EXPECT_EQ(result9->at("x"), x);
EXPECT_EQ(result9->at("y"), y);
}
TEST_F(TestIndexExpr, BoundSimplify) {
ir::Var S0 = ir::Var("S0");
ir::Var i = ir::Var(ir::Expr(0), ir::Expr(5), "i"); // i ∈ [0, 5)
ir::Var j = ir::Var(ir::Expr(0), S0, "j"); // j ∈ [0, S0)
ir::Expr q0 = i / Expr(5);
ir::Expr q1 = i / Expr(4);
ir::Expr q2 = i / Expr(6);
ir::Expr q3 = j / S0;
ir::Expr q4 = j / (S0 - 1);
ir::Expr q5 = j / (S0 + 1);
ir::Expr q6 = i % Expr(5);
ir::Expr q7 = i % Expr(4);
ir::Expr q8 = i % Expr(6);
ir::Expr q9 = j % S0;
ir::Expr q10 = j % (S0 - 1);
ir::Expr q11 = j % (S0 + 1);
EXPECT_EQ(q0.as_index().Normalize(ir::IndexExpr::OptLevel::kLevel3),
ir::Expr(0));
EXPECT_EQ(q1.as_index().Normalize(ir::IndexExpr::OptLevel::kLevel3),
i / Expr(4));
EXPECT_EQ(q2.as_index().Normalize(ir::IndexExpr::OptLevel::kLevel3),
ir::Expr(0));
EXPECT_EQ(q3.as_index().Normalize(ir::IndexExpr::OptLevel::kLevel3),
ir::Expr(0));
EXPECT_EQ(q4.as_index().Normalize(ir::IndexExpr::OptLevel::kLevel3),
j / (S0 + ir::Expr(-1)));
EXPECT_EQ(q5.as_index().Normalize(ir::IndexExpr::OptLevel::kLevel3),
ir::Expr(0));
EXPECT_EQ(q6.as_index().Normalize(ir::IndexExpr::OptLevel::kLevel3), i);
EXPECT_EQ(q7.as_index().Normalize(ir::IndexExpr::OptLevel::kLevel3),
i % Expr(4));
EXPECT_EQ(q8.as_index().Normalize(ir::IndexExpr::OptLevel::kLevel3), i);
EXPECT_EQ(q9.as_index().Normalize(ir::IndexExpr::OptLevel::kLevel3), j);
EXPECT_EQ(q10.as_index().Normalize(ir::IndexExpr::OptLevel::kLevel3),
j % (S0 + ir::Expr(-1)));
EXPECT_EQ(q11.as_index().Normalize(ir::IndexExpr::OptLevel::kLevel3), j);
}
} // namespace common
} // namespace cinn
+518
View File
@@ -0,0 +1,518 @@
// Copyright (c) 2024 CINN Authors. All Rights Reserved.
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
#include "paddle/cinn/common/iter_simplify.h"
#include <glog/logging.h>
#include <gtest/gtest.h>
#include "paddle/cinn/common/integer_set.h"
#include "paddle/cinn/common/ir_util.h"
#include "paddle/cinn/common/simplify_special_pattern.h"
#include "paddle/cinn/ir/op/ir_operators.h"
#include "paddle/cinn/ir/schedule/ir_schedule.h"
#include "paddle/cinn/ir/schedule/schedule_base.h"
namespace cinn {
namespace common {
#define ITER_MARK_VAR(var) \
ir::IterMark::Make(ir::IndexExpr(var.ptr()), var->upper_bound)
#define ITER_MARK_SUM(sum, ext) ir::IterMark::Make(sum, ext)
#define ITER_SPLIT(mark, ...) ir::IterSplit::Make(mark, ##__VA_ARGS__)
#define ITER_SUM(...) ir::IterSum::Make({__VA_ARGS__}, ir::IndexExpr(0))
#define ITER_SUM_WITH_BASE(base, ...) ir::IterSum::Make({__VA_ARGS__}, base)
#define TEST_EXPR(expr, expected, expr_norm) \
rewriter.Rewrite(&expr); \
EXPECT_EQ(expr, Expr(expected)); \
normalizer.Convert(&expr); \
EXPECT_EQ(expr, expr_norm);
class TestIterSimplify : public ::testing::Test {
public:
void SetUp() override {
i = ir::Var(ir::Expr(0), ir::Expr(2), "i").set_index(1);
j = ir::Var(ir::Expr(0), ir::Expr(4), "j").set_index(1);
k = ir::Var(ir::Expr(0), ir::Expr(8), "k").set_index(1);
i_j_k_fused =
ir::Var(ir::Expr(0), ir::Expr(64), "i_j_k_fused").set_index(1);
var_intervals = {
{"i", CasInterval(i->lower_bound, i->upper_bound - ir::Expr(1))},
{"j", CasInterval(j->lower_bound, j->upper_bound - ir::Expr(1))},
{"k", CasInterval(k->lower_bound, k->upper_bound - ir::Expr(1))},
{"i_j_k_fused",
CasInterval(i_j_k_fused->lower_bound,
i_j_k_fused->upper_bound - ir::Expr(1))}};
};
ir::Var i;
ir::Var j;
ir::Var k;
ir::Var i_j_k_fused;
cas_intervals_t var_intervals;
SymbolicExprAnalyzer analyzer{var_intervals};
};
TEST_F(TestIterSimplify, IterExprMake) {
// IterMark Make func.
auto mark_expr = ITER_MARK_VAR(i);
auto mark_expr_ = ITER_MARK_VAR(j);
// IterSplit Make func.
auto split_0_expr = ITER_SPLIT(mark_expr);
auto split_1_expr = ITER_SPLIT(mark_expr, ir::IndexExpr(1));
auto split_2_expr = ITER_SPLIT(
mark_expr, ir::IndexExpr(1), ir::IndexExpr(2), ir::IndexExpr(1));
auto split_3_expr = ITER_SPLIT(
mark_expr, ir::IndexExpr(2), ir::IndexExpr(2), ir::IndexExpr(1));
auto split_4_expr = ITER_SPLIT(
mark_expr_, ir::IndexExpr(1), ir::IndexExpr(2), ir::IndexExpr(1));
// IterSum Make func.
auto sum_expr = ITER_SUM(split_0_expr, split_1_expr, split_2_expr);
auto mark = mark_expr.As<ir::IterMark>();
auto split_0 = split_0_expr.As<ir::IterSplit>();
auto split_1 = split_1_expr.As<ir::IterSplit>();
auto split_2 = split_2_expr.As<ir::IterSplit>();
auto sum = sum_expr.As<ir::IterSum>();
EXPECT_EQ(mark->source, ir::IndexExpr(i.ptr()));
EXPECT_EQ(mark->extent, ir::IndexExpr(2));
EXPECT_EQ(split_0->source, mark_expr);
EXPECT_EQ(split_0->lower_factor, ir::IndexExpr(1));
EXPECT_EQ(split_0->extent, ir::IndexExpr(2));
EXPECT_EQ(split_0->scale, ir::IndexExpr(1));
EXPECT_EQ(split_1->source, mark_expr);
EXPECT_EQ(split_1->lower_factor, ir::IndexExpr(1));
EXPECT_EQ(split_1->extent, ir::IndexExpr(2));
EXPECT_EQ(split_1->scale, ir::IndexExpr(1));
EXPECT_EQ(split_2->source, mark_expr);
EXPECT_EQ(split_2->lower_factor, ir::IndexExpr(1));
EXPECT_EQ(split_2->extent, ir::IndexExpr(2));
EXPECT_EQ(split_2->scale, ir::IndexExpr(1));
EXPECT_EQ(sum->args.size(), 3);
EXPECT_EQ(sum->base, Expr(0));
EXPECT_NE(mark_expr, mark_expr_);
EXPECT_EQ(split_0_expr, split_1_expr);
EXPECT_EQ(split_1_expr, split_2_expr);
EXPECT_NE(split_2_expr, split_3_expr);
}
TEST_F(TestIterSimplify, conversion) {
IterMapRewriter rewriter{{i}, analyzer};
IterMapToExprNormalizer normalizer{analyzer};
ir::Expr e1 = i;
auto gt = ITER_SUM(ITER_SPLIT(ITER_MARK_VAR(i)));
TEST_EXPR(e1, gt, e1);
}
TEST_F(TestIterSimplify, add) {
IterMapRewriter rewriter{{i, j, k}, analyzer};
IterMapToExprNormalizer normalizer{analyzer};
auto gt1 =
ITER_SUM(ITER_SPLIT(ITER_MARK_VAR(i)), ITER_SPLIT(ITER_MARK_VAR(j)));
auto gt2 = ITER_SUM_WITH_BASE(ir::IndexExpr(5),
ITER_SPLIT(ITER_MARK_VAR(i)),
ITER_SPLIT(ITER_MARK_VAR(j)),
ITER_SPLIT(ITER_MARK_VAR(k)));
auto gt3 = ITER_SUM(ITER_SPLIT(ITER_MARK_VAR(i), ir::IndexExpr(2)));
auto gt4 = ITER_SUM_WITH_BASE(ir::IndexExpr(12));
ir::Expr e1 = i + j;
ir::Expr e2 = i + j + k + 5;
ir::Expr e3 = i + i;
ir::Expr e4 = Expr(7) + Expr(5);
TEST_EXPR(e1, gt1, i + j);
TEST_EXPR(e2, gt2, i + j + k + 5);
TEST_EXPR(e3, gt3, i * 2);
TEST_EXPR(e4, gt4, Expr(12));
}
TEST_F(TestIterSimplify, sub) {
IterMapRewriter rewriter{{i, j, k}, analyzer};
IterMapToExprNormalizer normalizer{analyzer};
auto gt1 = ITER_SUM(ITER_SPLIT(ITER_MARK_VAR(i)),
ITER_SPLIT(ITER_MARK_VAR(j), ir::IndexExpr(-1)));
auto gt2 =
ITER_SUM_WITH_BASE(ir::IndexExpr(5),
ITER_SPLIT(ITER_MARK_VAR(i)),
ITER_SPLIT(ITER_MARK_VAR(j)),
ITER_SPLIT(ITER_MARK_VAR(k), ir::IndexExpr(-1)));
auto gt3 = ITER_SUM(ITER_SPLIT(ITER_MARK_VAR(i), ir::IndexExpr(0)));
auto gt4 = ITER_SUM_WITH_BASE(ir::IndexExpr(2));
ir::Expr e1 = i - j;
ir::Expr e2 = i + j - k + 5;
ir::Expr e3 = i - i;
ir::Expr e4 = Expr(7) - Expr(5);
TEST_EXPR(e1, gt1, (j * -1) + i);
TEST_EXPR(e2, gt2, i + j + (k * -1) + 5);
TEST_EXPR(e3, gt3, Expr(0));
TEST_EXPR(e4, gt4, Expr(2));
}
TEST_F(TestIterSimplify, mul) {
IterMapRewriter rewriter{{i, j, k}, analyzer};
IterMapToExprNormalizer normalizer{analyzer};
auto gt1 = ITER_SUM(ITER_SPLIT(ITER_MARK_VAR(i), ir::IndexExpr(2)),
ITER_SPLIT(ITER_MARK_VAR(j)));
auto gt2 = ITER_SUM(ITER_SPLIT(ITER_MARK_VAR(i), ir::IndexExpr(2)),
ITER_SPLIT(ITER_MARK_VAR(j), ir::IndexExpr(2)),
ITER_SPLIT(ITER_MARK_VAR(k)));
auto gt3 = ITER_SUM_WITH_BASE(ir::IndexExpr(10),
ITER_SPLIT(ITER_MARK_VAR(i), ir::IndexExpr(2)),
ITER_SPLIT(ITER_MARK_VAR(j), ir::IndexExpr(2)),
ITER_SPLIT(ITER_MARK_VAR(k)));
auto gt4 = ITER_SUM_WITH_BASE(ir::IndexExpr(35));
ir::Expr e1 = i * 2 + j;
ir::Expr e2 = (i + j) * 2 + k;
ir::Expr e3 = (i + j + 5) * 2 + k;
ir::Expr e4 = Expr(7) * Expr(5);
TEST_EXPR(e1, gt1, i * 2 + j);
TEST_EXPR(e2, gt2, (i + j) * 2 + k);
TEST_EXPR(e3, gt3, (i + j) * 2 + k + 10);
TEST_EXPR(e4, gt4, Expr(35));
}
TEST_F(TestIterSimplify, div) {
IterMapRewriter rewriter{{i, j, k, i_j_k_fused}, analyzer};
IterMapToExprNormalizer normalizer{analyzer};
auto gt1 = ITER_SUM(ITER_SPLIT(ITER_MARK_VAR(i_j_k_fused),
ir::IndexExpr(8),
ir::IndexExpr(8),
ir::IndexExpr(1)));
auto gt2 = ITER_SUM(ITER_SPLIT(ITER_MARK_VAR(i_j_k_fused),
ir::IndexExpr(32),
ir::IndexExpr(2),
ir::IndexExpr(1)));
auto gt3 = ITER_SUM(ITER_SPLIT(ITER_MARK_VAR(i_j_k_fused)));
auto gt4 = ITER_SUM(ITER_SPLIT(ITER_MARK_VAR(i_j_k_fused), ir::IndexExpr(2)));
auto gt5 = ITER_SUM(ITER_SPLIT(ITER_MARK_VAR(i_j_k_fused),
ir::IndexExpr(2),
ir::IndexExpr(32),
ir::IndexExpr(1)));
auto gt6 = ITER_SUM(ITER_SPLIT(
ITER_MARK_SUM(ITER_SUM_WITH_BASE(ir::IndexExpr(8),
ITER_SPLIT(ITER_MARK_VAR(i_j_k_fused))),
ir::IndexExpr(72)),
ir::IndexExpr(16),
ir::IndexExpr(5),
ir::IndexExpr(1)));
auto gt7 = ITER_SUM(ITER_SPLIT(
ITER_MARK_SUM(ITER_SUM_WITH_BASE(ir::IndexExpr(1),
ITER_SPLIT(ITER_MARK_VAR(i_j_k_fused))),
ir::IndexExpr(65)),
ir::IndexExpr(2),
ir::IndexExpr(33),
ir::IndexExpr(1)));
auto gt8 = ITER_SUM_WITH_BASE(ir::IndexExpr(2),
ITER_SPLIT(ITER_MARK_VAR(i_j_k_fused),
ir::IndexExpr(8),
ir::IndexExpr(8),
ir::IndexExpr(1)));
auto gt9 = ITER_SUM_WITH_BASE(
ir::IndexExpr(2),
ITER_SPLIT(ITER_MARK_VAR(i_j_k_fused), ir::IndexExpr(2)));
auto gt10 = ITER_SUM(ITER_SPLIT(
ITER_MARK_SUM(ITER_SUM_WITH_BASE(ir::IndexExpr(1),
ITER_SPLIT(ITER_MARK_VAR(i_j_k_fused))),
ir::IndexExpr(65)),
ir::IndexExpr(8),
ir::IndexExpr(9),
ir::IndexExpr(1)));
auto gt11 = ITER_SUM_WITH_BASE(ir::IndexExpr(3));
auto gt12 = ITER_SUM_WITH_BASE(ir::IndexExpr(3));
auto gt13 = ITER_SUM_WITH_BASE(ir::IndexExpr(15));
auto gt14 = ITER_SUM_WITH_BASE(ir::IndexExpr(0));
ir::Expr e1 = i_j_k_fused / 8;
ir::Expr e2 = i_j_k_fused / 8 / 4;
ir::Expr e3 = i_j_k_fused / 1;
ir::Expr e4 = i_j_k_fused * 16 / 8;
ir::Expr e5 = i_j_k_fused * 8 / 16;
ir::Expr e6 = (i_j_k_fused + 8) / 16;
ir::Expr e7 = (i_j_k_fused * 8 + 8) / 16;
ir::Expr e8 = (i_j_k_fused + 16) / 8;
ir::Expr e9 = (i_j_k_fused * 16 + 16) / 8;
ir::Expr e10 = (i_j_k_fused + 1) / 8;
ir::Expr e11 = Expr(15) / Expr(5);
ir::Expr e12 = Expr(15) / Expr(4);
ir::Expr e13 = Expr(15) / Expr(1);
ir::Expr e14 = Expr(0) / Expr(4);
TEST_EXPR(e1, gt1, i_j_k_fused / 8);
TEST_EXPR(e2, gt2, i_j_k_fused / 32);
TEST_EXPR(e3, gt3, i_j_k_fused);
TEST_EXPR(e4, gt4, i_j_k_fused * 2);
TEST_EXPR(e5, gt5, i_j_k_fused / 2);
TEST_EXPR(e6, gt6, (i_j_k_fused + 8) / 16);
TEST_EXPR(e7, gt7, (i_j_k_fused + 1) / 2);
TEST_EXPR(e8, gt8, i_j_k_fused / 8 + 2);
TEST_EXPR(e9, gt9, i_j_k_fused * 2 + 2);
TEST_EXPR(e10, gt10, (i_j_k_fused + 1) / 8);
TEST_EXPR(e11, gt11, Expr(3));
TEST_EXPR(e12, gt12, Expr(3));
TEST_EXPR(e13, gt13, Expr(15));
TEST_EXPR(e14, gt14, Expr(0));
}
TEST_F(TestIterSimplify, mod) {
IterMapRewriter rewriter{{i, j, k, i_j_k_fused}, analyzer};
IterMapToExprNormalizer normalizer{analyzer};
auto gt1 = ITER_SUM(ITER_SPLIT(ITER_MARK_VAR(i_j_k_fused),
ir::IndexExpr(1),
ir::IndexExpr(8),
ir::IndexExpr(1)));
auto gt2 = ITER_SUM(ITER_SPLIT(ITER_MARK_VAR(i_j_k_fused),
ir::IndexExpr(8),
ir::IndexExpr(4),
ir::IndexExpr(1)));
auto gt3 = ITER_SUM_WITH_BASE(ir::IndexExpr(0));
auto gt4 = ITER_SUM_WITH_BASE(ir::IndexExpr(0));
auto gt5 = ITER_SUM(ITER_SPLIT(ITER_MARK_VAR(i_j_k_fused),
ir::IndexExpr(1),
ir::IndexExpr(2),
ir::IndexExpr(8)));
auto gt6 = ITER_SUM(ITER_SPLIT(
ITER_MARK_SUM(ITER_SUM_WITH_BASE(ir::IndexExpr(8),
ITER_SPLIT(ITER_MARK_VAR(i_j_k_fused))),
ir::IndexExpr(72)),
ir::IndexExpr(1),
ir::IndexExpr(16),
ir::IndexExpr(1)));
auto gt7 = ITER_SUM(ITER_SPLIT(
ITER_MARK_SUM(ITER_SUM_WITH_BASE(ir::IndexExpr(1),
ITER_SPLIT(ITER_MARK_VAR(i_j_k_fused),
ir::IndexExpr(1),
ir::IndexExpr(64),
ir::IndexExpr(1))),
ir::IndexExpr(65)),
ir::IndexExpr(1),
ir::IndexExpr(2),
ir::IndexExpr(8)));
auto gt8 = ITER_SUM(ITER_SPLIT(ITER_MARK_VAR(i_j_k_fused),
ir::IndexExpr(1),
ir::IndexExpr(8),
ir::IndexExpr(1)));
auto gt9 = ITER_SUM_WITH_BASE(ir::IndexExpr(0));
auto gt10 = ITER_SUM(ITER_SPLIT(
ITER_MARK_SUM(ITER_SUM_WITH_BASE(ir::IndexExpr(1),
ITER_SPLIT(ITER_MARK_VAR(i_j_k_fused))),
ir::IndexExpr(65)),
ir::IndexExpr(1),
ir::IndexExpr(8),
ir::IndexExpr(1)));
auto gt11 = ITER_SUM_WITH_BASE(ir::IndexExpr(0));
auto gt12 = ITER_SUM_WITH_BASE(ir::IndexExpr(3));
auto gt13 = ITER_SUM_WITH_BASE(ir::IndexExpr(0));
auto gt14 = ITER_SUM_WITH_BASE(ir::IndexExpr(0));
ir::Expr e1 = i_j_k_fused % 8;
ir::Expr e2 = i_j_k_fused / 8 % 4;
ir::Expr e3 = i_j_k_fused % 1;
ir::Expr e4 = i_j_k_fused * 16 % 8;
ir::Expr e5 = i_j_k_fused * 8 % 16;
ir::Expr e6 = (i_j_k_fused + 8) % 16;
ir::Expr e7 = (i_j_k_fused * 8 + 8) % 16;
ir::Expr e8 = (i_j_k_fused + 16) % 8;
ir::Expr e9 = (i_j_k_fused * 16 + 16) % 8;
ir::Expr e10 = (i_j_k_fused + 1) % 8;
ir::Expr e11 = Expr(15) % Expr(5);
ir::Expr e12 = Expr(15) % Expr(4);
ir::Expr e13 = Expr(15) % Expr(1);
ir::Expr e14 = Expr(0) % Expr(4);
TEST_EXPR(e1, gt1, i_j_k_fused % 8);
TEST_EXPR(e2, gt2, i_j_k_fused % 32 / 8);
TEST_EXPR(e3, gt3, Expr(0));
TEST_EXPR(e4, gt4, Expr(0));
TEST_EXPR(e5, gt5, i_j_k_fused % 2 * 8);
TEST_EXPR(e6, gt6, (i_j_k_fused + 8) % 16);
TEST_EXPR(e7, gt7, (i_j_k_fused + 1) % 2 * 8);
TEST_EXPR(e8, gt8, i_j_k_fused % 8);
TEST_EXPR(e9, gt9, Expr(0));
TEST_EXPR(e10, gt10, (i_j_k_fused + 1) % 8);
TEST_EXPR(e11, gt11, Expr(0));
TEST_EXPR(e12, gt12, Expr(3));
TEST_EXPR(e13, gt13, Expr(0));
TEST_EXPR(e14, gt14, Expr(0));
}
TEST_F(TestIterSimplify, fuse_not_same_source) {
IterMapRewriter rewriter{{i, j, k, i_j_k_fused}, analyzer};
IterMapToExprNormalizer normalizer{analyzer};
auto gt1 = ITER_SUM(ITER_SPLIT(
ITER_MARK_SUM(ITER_SUM(ITER_SPLIT(ITER_MARK_VAR(i), ir::IndexExpr(32)),
ITER_SPLIT(ITER_MARK_VAR(j), ir::IndexExpr(8)),
ITER_SPLIT(ITER_MARK_VAR(k), ir::IndexExpr(1))),
ir::IndexExpr(64)),
ir::IndexExpr(8),
ir::IndexExpr(8),
ir::IndexExpr(1)));
ir::Expr e1 = (i * 32 + j * 8 + k) / 8;
ir::Expr e2 = (i * 32 + j * 7) / 8;
TEST_EXPR(e1, gt1, (i * 32 + j * 8 + k) / 8);
EXPECT_ANY_THROW(rewriter.Rewrite(&e2));
}
TEST_F(TestIterSimplify, fuse_same_source) {
IterMapRewriter rewriter{{i, j, k, i_j_k_fused}, analyzer};
IterMapToExprNormalizer normalizer{analyzer};
auto gt1 = ITER_SUM(ITER_SPLIT(ITER_MARK_VAR(i_j_k_fused),
ir::IndexExpr(32),
ir::IndexExpr(2),
ir::IndexExpr(1)));
auto gt2 = ITER_SUM(ITER_SPLIT(ITER_MARK_VAR(i_j_k_fused),
ir::IndexExpr(8),
ir::IndexExpr(4),
ir::IndexExpr(1)));
auto gt3 = ITER_SUM(ITER_SPLIT(ITER_MARK_VAR(i_j_k_fused),
ir::IndexExpr(1),
ir::IndexExpr(8),
ir::IndexExpr(1)));
auto gt4 = ITER_SUM(ITER_SPLIT(ITER_MARK_VAR(i_j_k_fused),
ir::IndexExpr(32),
ir::IndexExpr(2),
ir::IndexExpr(1)),
ITER_SPLIT(ITER_MARK_VAR(i_j_k_fused),
ir::IndexExpr(32),
ir::IndexExpr(1),
ir::IndexExpr(1)));
ir::Expr e1 = (i_j_k_fused / 16 / 2 * 32 + i_j_k_fused / 16 % 2 * 16 +
i_j_k_fused % 16) /
8 / 4;
ir::Expr e2 =
(i_j_k_fused / 32 * 32 + i_j_k_fused / 16 % 2 * 16 + i_j_k_fused % 16) /
8 % 4;
ir::Expr e3 =
(i_j_k_fused / 32 * 32 + i_j_k_fused / 16 % 2 * 16 + i_j_k_fused % 16) %
8;
ir::Expr e4 =
((i_j_k_fused / 16) / 2) +
((((i_j_k_fused % 16) / 8) + (2 * ((i_j_k_fused / 16) % 2))) / 4);
ir::Expr e5 = (((i_j_k_fused % 16) / 8) + ((4 * ((i_j_k_fused / 16) / 2)) +
(2 * ((i_j_k_fused / 16) % 2)))) %
4;
ir::Expr e6 = ((i_j_k_fused % 16) + ((32 * ((i_j_k_fused / 16) / 2)) +
(16 * ((i_j_k_fused / 16) % 2)))) %
8;
TEST_EXPR(e1, gt1, i_j_k_fused / 32);
TEST_EXPR(e2, gt2, i_j_k_fused % 32 / 8);
TEST_EXPR(e3, gt3, i_j_k_fused % 8);
TEST_EXPR(e4, gt4, i_j_k_fused / 32);
TEST_EXPR(e5, gt2, i_j_k_fused % 32 / 8);
TEST_EXPR(e6, gt3, i_j_k_fused % 8);
}
TEST_F(TestIterSimplify, SimplifyBindings) {
std::vector<ir::Var> block_vars;
std::vector<ir::Expr> iter_values;
std::vector<ir::Expr> shape = {Expr(2), Expr(4), Expr(8)};
std::vector<Var> axis_vars = cinn::common::GenDefaultAxis(3);
// Create block vars and axis vars
for (int i = 0; i < shape.size(); ++i) {
block_vars.push_back(ir::Var(Expr(0),
shape[i],
cinn::UniqName("b" + std::to_string(i)),
false,
false)
.set_index(1));
axis_vars[i]->is_reduce_axis = false;
iter_values.push_back(axis_vars[i]);
}
// Create ScheduleBlock body
ir::Expr body_ = ir::ScheduleBlockRealize::Make(
iter_values,
ir::ScheduleBlock::Make(block_vars, {}, {}, "Test", Expr(0)));
// Create For loops
auto body = body_;
for (int i = shape.size() - 1; i >= 0; --i) {
ir::Var loop_var = axis_vars[i];
ir::Expr loop_extent = shape[i];
body = ir::For::Make(loop_var,
Expr(0),
loop_extent,
ir::ForType::Serial,
ir::DeviceAPI::Host,
ir::Block::Make({body}));
}
// Create outer ScheduleBlockRealize
ir::Expr body_outer = ir::ScheduleBlockRealize::Make(
{}, ir::ScheduleBlock::Make({}, {}, {}, "test1", body));
// Create ir schedule
ir::ModuleExpr mod_expr({ir::Block::Make({body_outer})});
ir::IRSchedule ir_sch(mod_expr);
std::vector<ir::Expr> loops = ir_sch.GetLoops(body_);
// Apply Fuse and Split
ir::Expr loop_fuse = ir_sch.Fuse(loops);
std::vector<ir::Expr> loops_split = ir_sch.Split(loop_fuse, {2, 2, 16});
ir::Expr loop_fuse_2 = ir_sch.Fuse(loops_split);
// Apply SimplifyBindings
SimplifyBlockBinding::SimplifyBindings(loop_fuse_2, {}, analyzer);
// Check result
auto for_op = loop_fuse_2.As<ir::For>();
auto simplified_values = for_op->body.As<ir::Block>()
->stmts[0]
.As<ir::ScheduleBlockRealize>()
->iter_values;
auto f = for_op->loop_var;
EXPECT_EQ(simplified_values[0], f / 32);
EXPECT_EQ(simplified_values[1], f % 32 / 8);
EXPECT_EQ(simplified_values[2], f % 8);
}
TEST_F(TestIterSimplify, MergeMulMod) {
auto S0 = ir::Var(ir::Expr(0), ir::Expr(4), "S0").set_index(1);
auto S1 = ir::Var(ir::Expr(0), ir::Expr(256), "S1").set_index(1);
auto S2 = ir::Var(ir::Expr(0), ir::Expr(13), "S2").set_index(1);
auto e1 = ((((((((S0 * 256) + S1) + (S2 * 1024)) / 2500) * 50) +
(((((S0 * 256) + S1) + (S2 * 1024)) % 2500) / 50)) *
50) +
((((S0 * 256) + S1) + (S2 * 1024)) % 50));
auto e2 = ((((((S0 * 256) + S1) + (S2 * 1024)) / 2500) + -4) * 2500) +
((((S0 * 256) + S1) + (S2 * 1024)) % 2500);
auto e3 = (S1 / 784 * 28 + S1 % 784 / 28) * 28 + S1 % 28;
EXPECT_EQ(MergeMulMod(e1), (((S0 * 256) + S1) + (S2 * 1024)));
EXPECT_EQ(MergeMulMod(e2), ((((S0 * 256) + S1) + (S2 * 1024)) + -10000));
EXPECT_EQ(MergeMulMod(e3), S1);
}
} // namespace common
} // namespace cinn
+119
View File
@@ -0,0 +1,119 @@
// Copyright (c) 2023 PaddlePaddle Authors. All Rights Reserved.
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
#include <glog/logging.h>
#include <gtest/gtest.h>
#include <sstream>
#include "paddle/cinn/adt/generate_map_expr.h"
#include "paddle/cinn/adt/map_expr_ctx.h"
#include "paddle/cinn/adt/print.h"
#include "paddle/cinn/hlir/dialect/operator/ir/cinn_op.h"
#include "paddle/cinn/hlir/dialect/operator/ir/manual_op.h"
#include "paddle/cinn/hlir/dialect/operator/ir/op_dialect.h"
#include "paddle/cinn/runtime/flags.h"
#include "paddle/fluid/pir/dialect/operator/ir/op_dialect.h"
#include "paddle/fluid/pir/dialect/operator/ir/pd_op.h"
#include "paddle/pir/include/core/ir_context.h"
#include "paddle/pir/include/core/program.h"
#include "paddle/pir/include/dialect/shape/ir/shape_dialect.h"
#include "paddle/pir/include/dialect/shape/ir/shape_op.h"
#include "paddle/pir/include/dialect/shape/transforms/shape_optimization_pass.h"
#include "paddle/pir/include/pass/pass_manager.h"
#include "test/cpp/pir/tools/test_pir_utils.h"
PD_DECLARE_bool(cinn_enable_map_expr);
PD_DECLARE_bool(cinn_enable_map_expr_dynamic_shape);
PD_DECLARE_bool(cinn_enable_map_expr_index_detail);
namespace {
std::string Trim(const std::string& doc) {
std::stringstream oss{doc};
std::stringstream ret{};
std::string str;
while (oss >> str) {
std::size_t size = str.size();
for (; size > 0 && str[size - 1] == ' '; --size) {
}
ret << str.substr(0, size) << std::endl;
}
return ret.str();
}
} // namespace
TEST(MapExpr, ElementWise_Fusion_0) {
cinn::adt::UniqueId::ResetSeqNumber(0);
::pir::IrContext* ctx = ::pir::IrContext::Instance();
::pir::Program program(ctx);
ctx->GetOrRegisterDialect<pir::shape::ShapeDialect>();
ctx->GetOrRegisterDialect<paddle::dialect::OperatorDialect>();
ctx->GetOrRegisterDialect<cinn::dialect::OperatorDialect>();
phi::DDim dims_D_2 = {-1, 1};
pir::Value value1 =
test::CreateDenseTensorOp(ctx, dims_D_2, {"op1_attr"}, {"op1_name"})
->result(0);
pir::Value value2 =
test::CreateDenseTensorOp(ctx, dims_D_2, {"op2_attr"}, {"op2_name"})
->result(0);
::pir::Builder builder = ::pir::Builder(ctx, program.block());
builder.Build<paddle::dialect::SubtractOp>(
value1, builder.Build<paddle::dialect::ExpOp>(value2).result(0));
::pir::PassManager pass_manager(ctx);
// TODO(@jiahy0825): use CreateShapeOptimizationPass() instead of
// CreateInferSymbolicShapePass() which is a fake pass
/*
pass_manager.AddPass(::pir::CreateInferSymbolicShapePass(shape_analysis));
pass_manager.Run(&program);
std::vector<pir::Operation*> vec_op;
for (auto& op : *program.block()) {
vec_op.push_back(&op);
}
auto res = cinn::dialect::ir::OpFusionPassInternal(vec_op);
ASSERT_EQ(res.size(), 1u);
ASSERT_EQ(res[0]->ops.size(), program.block()->size());
auto group_list = cinn::dialect::ir::GeneralFusionMergePassInternal(res);
ASSERT_EQ(group_list.size(), 1u);
FLAGS_cinn_enable_map_expr = true;
FLAGS_cinn_enable_map_expr_dynamic_shape = true;
FLAGS_cinn_enable_map_expr_index_detail = true;
auto group = group_list.at(0);
group->shape_analysis = shape_analysis;
cinn::adt::TryGenerateMapExprFromGroup(group);
std::string map_expr_str =
cinn::adt::ToTxtString(group->map_expr_ctx().map_expr(), "MapExprTest");
std::string target_str = R"TEST(
MapExprTest(t_var_2, t_var_1) {
AnchoredMapStmt(t_var_0) {
MapStmt([i_59, i_60]) {
exp(
&t_var[IndexDot([BI(i_59, sym_17), 0], [sym_17, 1])],
t_var_1[IndexDot([BI(i_59, sym_17), 0], [sym_17, 1])]);
subtract(
&t_var_0[IndexDot([i_59, i_60], [sym_17, 1])],
t_var_2[IndexDot([BI(i_59, sym_17), 0], [sym_17, 1])],
t_var[IndexDot([BI(i_59, sym_17), 0], [sym_17, 1])]);
}
}
}
)TEST";
ASSERT_EQ(Trim(map_expr_str), Trim(target_str));
*/
}
@@ -0,0 +1,144 @@
// Copyright (c) 2024 PaddlePaddle Authors. All Rights Reserved.
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
#include <glog/logging.h>
#include <gtest/gtest.h>
#include <sstream>
#include "paddle/cinn/ir/op/ir_operators.h"
#include "paddle/cinn/optim/ir_simplify.h"
#include "paddle/cinn/optim/merge_block_utils.h"
namespace cinn {
namespace optim {
namespace {
bool IsBlockForAllEqual(const ForTreeNode& first, const ForTreeNode& second) {
auto ForVarExtentEqual = [&](const ForTreeNode& first,
const ForTreeNode& second) -> bool {
const ir::Expr lhs = first.val->extent();
const ir::Expr rhs = second.val->extent();
if (lhs != rhs) {
return false;
}
return true;
};
if (!ForVarExtentEqual(first, second)) return false;
if (first.children.size() != second.children.size()) return false;
for (size_t i = 0; i < first.children.size(); ++i) {
if (!IsBlockForAllEqual(first.children[i], second.children[i])) {
return false;
}
}
return true;
}
ir::stmt::For MakeForLoops(const std::vector<int> extents, int index) {
ir::stmt::StmtRef body_stmt;
if (index == extents.size() - 1) {
body_stmt = ir::stmt::Schedule(
std::vector<Var>(),
std::vector<Expr>(),
std::vector<Expr>(),
std::vector<Expr>(),
"block",
ir::stmt::BlockRef(std::vector<ir::stmt::StmtRef>()));
} else {
body_stmt = MakeForLoops(extents, index + 1);
}
std::vector<ir::stmt::StmtRef> body = {body_stmt};
return ir::stmt::For(ir::Var("i"),
ir::Expr(0),
ir::Expr(extents[index]),
ir::ForType::Serial,
ir::DeviceAPI::CUDA,
ir::stmt::BlockRef(body),
ir::VectorizeInfo(),
ir::BindInfo());
}
void TestHelper(const std::vector<int>& extents1,
const std::vector<int>& extents2,
bool is_same) {
auto for_loop1 = MakeForLoops(extents1, 0);
auto for_loop2 = MakeForLoops(extents2, 0);
if (is_same) {
EXPECT_TRUE(CanMergeBlocks(for_loop1, for_loop2, IsBlockForAllEqual));
} else {
EXPECT_FALSE(CanMergeBlocks(for_loop1, for_loop2, IsBlockForAllEqual));
}
}
void TestHelper2(const std::vector<std::vector<int>>& extents1,
const std::vector<std::vector<int>>& extents2,
bool is_same) {
auto MakeNestLoops =
[&](const std::vector<std::vector<int>>& extents) -> ir::stmt::For {
std::vector<ir::stmt::StmtRef> for_loops;
for (size_t i = 0; i < extents.size(); ++i) {
for_loops.push_back(MakeForLoops(extents[i], 0));
}
ir::stmt::BlockRef block(for_loops);
ir::stmt::For for_stmt = ir::stmt::For(ir::Var("i"),
ir::Expr(0),
ir::Expr(1),
ir::ForType::Serial,
ir::DeviceAPI::CUDA,
block,
ir::VectorizeInfo(),
ir::BindInfo());
return for_stmt;
};
auto for_stmt1 = MakeNestLoops(extents1);
auto for_stmt2 = MakeNestLoops(extents2);
if (is_same) {
EXPECT_TRUE(CanMergeBlocks(for_stmt1, for_stmt2, IsBlockForAllEqual));
} else {
EXPECT_FALSE(CanMergeBlocks(for_stmt1, for_stmt2, IsBlockForAllEqual));
}
}
TEST(ForInfo, ForInfoEqual) {
TestHelper({10}, {10}, true);
TestHelper({10, 5}, {10, 5}, true);
TestHelper({10, 5, 3}, {10, 5, 3}, true);
TestHelper2({{10}, {10}}, {{10}, {10}}, true);
TestHelper2({{10, 5}, {4, 7}}, {{10, 5}, {4, 7}}, true);
TestHelper2(
{{10, 5, 3}, {4, 7, 9}, {2, 8}}, {{10, 5, 3}, {4, 7, 9}, {2, 8}}, true);
}
TEST(ForInfo, ForInfoNotEqual) {
TestHelper({10}, {9}, false);
TestHelper({10, 5}, {10, 4}, false);
TestHelper({10, 5, 3}, {10, 5, 2}, false);
TestHelper2({{10}, {10}}, {{10}, {9}}, false);
TestHelper2({{10, 5}, {4, 7}}, {{10, 5}, {4, 3}}, false);
TestHelper2(
{{10, 5, 3}, {4, 7, 9}, {2, 8}}, {{10, 5, 3}, {4, 7, 9}, {2, 7}}, false);
}
} // namespace
} // namespace optim
} // namespace cinn