Skip to content

Commit 679f664

Browse files
DerrickYLJspectrometerHBH
authored andcommitted
[Layout] Fuse layouts w/ or w/t device tree (apache#48)
* Fuse layouts w/ or w/t device tree * fix comments except unit removal * delete cnt & remoce dev in attr
1 parent c3e0d03 commit 679f664

2 files changed

Lines changed: 498 additions & 6 deletions

File tree

‎src/tir/ir/layout.cc‎

Lines changed: 292 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -546,6 +546,7 @@ TileLayout::SplitMap TileLayout::GetSplitMap(
546546
}
547547

548548
/******** Normalization ********/
549+
549550
using NodeSet = std::unordered_set<IterTreeBase, ObjectPtrHash, ObjectPtrEqual>;
550551

551552
class DeepCopyMutator : public IterTreeMutator {
@@ -685,10 +686,11 @@ class UnitIterRemover : public IterTreeMutator {
685686
ICHECK(split_map_ != nullptr) << "InternalError: The split map should be defined";
686687
auto dev_it = split_map_->find(root);
687688
if (dev_it != split_map_->end()) {
688-
attr_map_->erase(dev_it->second);
689689
split_map_->erase(dev_it);
690690
}
691691
}
692+
} else {
693+
attr_map_->erase(root);
692694
}
693695
// remove this iter
694696
return NullOpt;
@@ -720,6 +722,9 @@ class UnitIterRemover : public IterTreeMutator {
720722
auto unit = IterTreeSplit(1, {});
721723
if (is_data_) {
722724
coeff_map_->insert({unit, 1});
725+
auto* n = root.CopyOnWrite();
726+
n->children = {unit};
727+
return GetRef<IterTreeBase>(n);
723728
} else {
724729
attr_map_->insert({unit, DeviceIterAttr::Replicate()});
725730
}
@@ -760,6 +765,7 @@ TileLayout RemoveUnitIter(TileLayout layout) {
760765
DataIterTree::CoeffMap coeff_map = data_tree.GetCoeffMap(data_leaves);
761766
if (!layout->device_tree.defined()) {
762767
// Only data tree is defined
768+
763769
auto new_data_root = UnitIterRemover::RemoveUnitIter(data_tree->root, true, &coeff_map);
764770
ICHECK(new_data_root.defined()) << "InternalError: The data tree should be defined";
765771
ICHECK_GT(new_data_root.value()->children.size(), 0)
@@ -783,9 +789,293 @@ TileLayout RemoveUnitIter(TileLayout layout) {
783789
}
784790
}
785791

792+
class EquivIterFuser : public IterTreeMutator {
793+
public:
794+
explicit EquivIterFuser(bool only_data, DataIterTree::CoeffMap* coeff_map,
795+
DeviceIterTree::AttrMap* attr_map = nullptr,
796+
TileLayout::SplitMap* split_map = nullptr,
797+
const DeviceIterTree* dev_tree = nullptr)
798+
: only_data_(only_data),
799+
coeff_map_(coeff_map),
800+
attr_map_(attr_map),
801+
split_map_(split_map),
802+
dev_tree_(dev_tree),
803+
ana_() {}
804+
805+
static Optional<IterTreeSplit> FuseEquivIter(IterTreeBase root, bool only_data,
806+
DataIterTree::CoeffMap* coeff_map,
807+
DeviceIterTree::AttrMap* attr_map = nullptr,
808+
TileLayout::SplitMap* split_map = nullptr,
809+
const DeviceIterTree* dev_tree = nullptr) {
810+
auto res = EquivIterFuser(only_data, coeff_map, attr_map, split_map, dev_tree).Visit(root);
811+
return res.as<IterTreeSplit>();
812+
}
813+
814+
PrimExpr ComputeGCD(PrimExpr a, PrimExpr b, arith::Analyzer* ana) {
815+
while (!ana->CanProveEqual(b, 0)) {
816+
PrimExpr temp = b;
817+
b = floormod(a, b);
818+
a = temp;
819+
}
820+
return a;
821+
}
822+
823+
bool CheckEquiv(PrimExpr a, PrimExpr b, arith::Analyzer* ana) { return ana->CanProveEqual(a, b); }
824+
825+
IterTreeSplit* Check_device_adjacent(const IterTreeBase& device1, const IterTreeBase& device2,
826+
IterTreeSplit* dev_node) {
827+
if (dev_node->IsLeaf()) {
828+
return nullptr;
829+
}
830+
// add recursive check for device leaf nodes
831+
tvm::tir::IterTreeSplit* curr_res = nullptr;
832+
std::vector<IterTreeBase> dev_children;
833+
for (const auto& child : (*dev_node)->children) {
834+
// recursively check all device leaf nodes
835+
auto* child_node = child.as<IterTreeSplitNode>();
836+
auto input_child_node = GetRef<IterTreeSplit>(child_node);
837+
auto new_curr_res = Check_device_adjacent(device1, device2, &input_child_node);
838+
ICHECK(!(curr_res != nullptr && new_curr_res != nullptr))
839+
<< "InternalError: device nodes can only be fused once!";
840+
curr_res = new_curr_res == nullptr ? curr_res : new_curr_res;
841+
dev_children.push_back(((Optional<IterTreeBase>)input_child_node).value());
842+
}
843+
if (curr_res == nullptr) {
844+
for (int i = (int)dev_children.size() - 1; i >= 1; --i) {
845+
const auto& dev_child1 = Downcast<IterTreeSplit>(dev_children[i]);
846+
const auto& dev_child2 = Downcast<IterTreeSplit>(dev_children[i - 1]);
847+
if (dev_child1.defined() && dev_child2.defined() && dev_child1.IsLeaf() &&
848+
dev_child2.IsLeaf() && dev_child1.same_as(device1) && dev_child2.same_as(device2)) {
849+
// Fuse device 1 and device 2 to one device leaf node with the new extent equal to
850+
// the product of the two
851+
auto new_extent = dev_child1->extent * dev_child2->extent;
852+
curr_res = new IterTreeSplit();
853+
*curr_res = IterTreeSplit(new_extent, {});
854+
// Replace the old dev_child1 and dev_child2 with *curr_res
855+
dev_children.erase(dev_children.begin() + i);
856+
dev_children[i - 1] = GetRef<IterTreeBase>((*curr_res).get());
857+
auto* n = (*dev_node).CopyOnWrite();
858+
n->children = dev_children;
859+
return curr_res;
860+
}
861+
}
862+
} else {
863+
// If changed at the bottom, update children at the top recursively
864+
auto* n = (*dev_node).CopyOnWrite();
865+
n->children = dev_children;
866+
return curr_res;
867+
}
868+
return curr_res;
869+
}
870+
871+
Optional<IterTreeBase> VisitIterSplit(const IterTreeSplitNode* node) final {
872+
auto root = GetRef<IterTreeSplit>(node);
873+
if (only_data_) {
874+
if (root.IsLeaf()) {
875+
// if it's a leaf, keep it
876+
return root;
877+
} else {
878+
// if it's a node, use bottom-up recursion
879+
std::vector<IterTreeBase> new_children;
880+
bool all_leaves = true;
881+
bool changed = false;
882+
bool fused = false;
883+
for (const auto& child : root->children) {
884+
// only normalize consider all-children-leaves case
885+
bool is_root_cur = is_root_;
886+
is_root_ = false;
887+
auto new_child = Visit(child);
888+
is_root_ = is_root_cur;
889+
auto new_child_node = GetRef<IterTreeSplit>(new_child.as<IterTreeSplitNode>());
890+
ICHECK(new_child.defined()) << "InternalError: New child must be valid";
891+
all_leaves = new_child_node.IsLeaf() ? all_leaves : false;
892+
new_children.push_back(new_child.value());
893+
changed |= (!new_child.same_as(child));
894+
}
895+
896+
PrimExpr gcd_coeff;
897+
PrimExpr min_coeff;
898+
if (all_leaves) {
899+
// compute GCD of coefficients of all leaves
900+
auto it = coeff_map_->find(Downcast<IterTreeSplit>(new_children[0]));
901+
ICHECK(it != coeff_map_->end()) << "InternalError: Coefficient not found for child";
902+
gcd_coeff = it->second;
903+
min_coeff = it->second;
904+
for (size_t i = 1; i < new_children.size(); ++i) {
905+
it = coeff_map_->find(Downcast<IterTreeSplit>(new_children[i]));
906+
ICHECK(it != coeff_map_->end()) << "InternalError: Coefficient not found for child";
907+
gcd_coeff = ComputeGCD(gcd_coeff, it->second, &ana_);
908+
min_coeff = tvm::min(it->second, min_coeff);
909+
}
910+
// check if normalize requirement is satisfied by coeff and extents
911+
fused = CheckEquiv(min_coeff, gcd_coeff, &ana_);
912+
PrimExpr divider_product = gcd_coeff;
913+
for (int i = (int)new_children.size() - 1; i >= 0; --i) {
914+
// reversal order to check mod and coeff
915+
const auto& child = Downcast<IterTreeSplit>(new_children[i]);
916+
auto coeff_it = coeff_map_->find(child);
917+
ICHECK(coeff_it != coeff_map_->end())
918+
<< "InternalError: Coefficient not found for child";
919+
if (!ana_.CanProveEqual(divider_product, coeff_it->second)) {
920+
fused = false;
921+
break;
922+
}
923+
divider_product *= Downcast<IterTreeSplit>(new_children[i])->extent;
924+
}
925+
}
926+
if (fused) {
927+
// fuse all children into the current node
928+
if (!is_root_) {
929+
for (int i = (int)new_children.size() - 1; i >= 0; --i) {
930+
coeff_map_->erase(GetRef<IterTreeSplit>(new_children[i].as<IterTreeSplitNode>()));
931+
}
932+
coeff_map_->erase(root);
933+
auto* n = root.CopyOnWrite();
934+
n->children = {};
935+
coeff_map_->insert({GetRef<IterTreeSplit>(n), gcd_coeff});
936+
return GetRef<IterTreeBase>(n);
937+
} else {
938+
// add a dummy leaf if normalizing root
939+
for (int i = (int)new_children.size() - 1; i >= 0; --i) {
940+
coeff_map_->erase(GetRef<IterTreeSplit>(new_children[i].as<IterTreeSplitNode>()));
941+
}
942+
coeff_map_->erase(root);
943+
auto* n = root.CopyOnWrite();
944+
auto unit = IterTreeSplit(root->extent, {});
945+
coeff_map_->insert({unit, gcd_coeff});
946+
n->children = {unit};
947+
return GetRef<IterTreeBase>(n);
948+
}
949+
} else {
950+
// no fuse happened
951+
if (changed) {
952+
auto* n = root.CopyOnWrite();
953+
n->children = new_children;
954+
return GetRef<IterTreeBase>(n);
955+
} else {
956+
return root;
957+
}
958+
}
959+
}
960+
} else {
961+
// Fuse the cases when both device tree and data tree are defined as follows:
962+
// For any two of data leaf nodes that share the same parent and adjacent to each other,
963+
// If two data nodes are mapped to two device nodes that are adjacent to each other,
964+
// then fuse those two data leaf nodes to one single data node with the extent of the new node
965+
// equal to the product of the extents of the two rest of data tree structure unchanged, and
966+
// fuse those two device nodes to one single node with the value being the product of those
967+
// two device nodes
968+
if (root.IsLeaf()) {
969+
// if it's a leaf, keep it
970+
return root;
971+
} else {
972+
// if it's a node, use bottom-up recursion
973+
std::vector<IterTreeBase> new_children;
974+
bool changed = false;
975+
for (const auto& child : root->children) {
976+
// only normalize consider all-children-leaves case
977+
bool is_root_cur = is_root_;
978+
is_root_ = false;
979+
auto new_child = Visit(child);
980+
is_root_ = is_root_cur;
981+
auto new_child_node = GetRef<IterTreeSplit>(new_child.as<IterTreeSplitNode>());
982+
ICHECK(new_child.defined()) << "InternalError: New child must be valid";
983+
new_children.push_back(new_child.value());
984+
changed |= (!new_child.same_as(child));
985+
}
986+
int children_cnt = (int)new_children.size();
987+
for (int i = children_cnt - 1; i >= 1; --i) {
988+
// reversal order to check mod and coeff
989+
const auto& child1 = Downcast<IterTreeSplit>(new_children[i]);
990+
const auto& child2 = Downcast<IterTreeSplit>(new_children[i - 1]);
991+
auto device1 = split_map_->find(child1);
992+
auto device2 = split_map_->find(child2);
993+
994+
if (device1 != split_map_->end() && device2 != split_map_->end()) {
995+
const tvm::tir::IterTreeSplit* fused_dev_node =
996+
Check_device_adjacent(device1->second, device2->second,
997+
const_cast<tvm::tir::IterTreeSplit*>(&(*dev_tree_)->root));
998+
if (fused_dev_node != nullptr) {
999+
// TODO: fuse child1 and child2 in data tree to a new single node, modify coeff,
1000+
// attr_map, split_map accordingly based on fused_dev_node
1001+
// Fuse child1 and child2 in data tree to a new single node
1002+
ICHECK(children_cnt >= 2)
1003+
<< "InterError: Data nodes cannot fall below zero after fusing";
1004+
PrimExpr new_extent = child1->extent * child2->extent;
1005+
IterTreeSplit new_data_node = IterTreeSplit(new_extent, {});
1006+
IterTreeSplit unit_data_node = IterTreeSplit(1, {});
1007+
// Update coeff_map
1008+
coeff_map_->erase(child1);
1009+
coeff_map_->erase(child2);
1010+
coeff_map_->insert({new_data_node, -1});
1011+
coeff_map_->insert({unit_data_node, -1});
1012+
// Update attr_map
1013+
attr_map_->erase(device1->second);
1014+
attr_map_->erase(device2->second);
1015+
attr_map_->insert({*fused_dev_node, DeviceIterAttr::Split(new_extent)});
1016+
// Update split_map
1017+
split_map_->erase(child1);
1018+
split_map_->erase(child2);
1019+
split_map_->insert({new_data_node, *fused_dev_node});
1020+
// Replace the old children with the new fused node
1021+
// new_children.erase(new_children.begin() + i);
1022+
new_children[i - 1] = new_data_node;
1023+
new_children[i] = unit_data_node;
1024+
children_cnt -= 1;
1025+
changed = true;
1026+
}
1027+
}
1028+
}
1029+
if (changed) {
1030+
auto* n = root.CopyOnWrite();
1031+
n->children = new_children;
1032+
return GetRef<IterTreeBase>(n);
1033+
} else {
1034+
return root;
1035+
}
1036+
}
1037+
}
1038+
}
1039+
1040+
private:
1041+
bool only_data_; // if true, only data tree is defined;
1042+
// otherwise, both data&device trees are defined
1043+
bool is_root_{true};
1044+
DataIterTree::CoeffMap* coeff_map_;
1045+
DeviceIterTree::AttrMap* attr_map_;
1046+
TileLayout::SplitMap* split_map_;
1047+
const DeviceIterTree* dev_tree_;
1048+
arith::Analyzer ana_; // Analyzer to simplify and check expressions
1049+
};
1050+
1051+
TileLayout FuseEquivIter(TileLayout layout) {
1052+
const auto& data_tree = layout->data_tree;
1053+
const auto& data_leaves = data_tree.GetLeaves();
1054+
DataIterTree::CoeffMap coeff_map = data_tree.GetCoeffMap(data_leaves);
1055+
if (!layout->device_tree.defined()) {
1056+
auto new_data_root = EquivIterFuser::FuseEquivIter(data_tree->root, true, &coeff_map);
1057+
ICHECK(new_data_root.defined()) << "InternalError: The data tree should be defined";
1058+
return TileLayout::FromMaps(new_data_root.value(), NullOpt, coeff_map, {}, {});
1059+
} else {
1060+
// Both data tree and device tree are defined
1061+
const auto& dev_tree = layout->device_tree.value();
1062+
const auto& dev_leaves = dev_tree.GetLeaves();
1063+
DeviceIterTree::AttrMap attr_map = dev_tree.GetAttrMap(dev_leaves);
1064+
TileLayout::SplitMap split_map = layout.GetSplitMap(data_leaves, dev_leaves);
1065+
auto new_data_root = EquivIterFuser::FuseEquivIter(data_tree->root, false, &coeff_map,
1066+
&attr_map, &split_map, &dev_tree);
1067+
// auto new_dev_root = GetRef<IterTreeSplit>(dev_tree->root);
1068+
ICHECK(new_data_root.defined() && dev_tree->root.defined())
1069+
<< "InternalError: The data tree and device tree should be both defined";
1070+
return RemoveUnitIter(TileLayout::FromMaps(new_data_root.value(), dev_tree->root, coeff_map,
1071+
attr_map, split_map, layout->from, layout->to));
1072+
}
1073+
}
1074+
7861075
TileLayout NormalizeTileLayout(TileLayout layout) {
7871076
TileLayout res = Downcast<TileLayout>(Deduplicate({layout})[0]);
7881077
res = RemoveUnitIter(res);
1078+
res = FuseEquivIter(res);
7891079
return std::move(res);
7901080
}
7911081

@@ -1015,4 +1305,5 @@ PrimExpr SwizzleLayoutNode::Apply(const Array<PrimExpr>& coord) const {
10151305
}
10161306

10171307
} // namespace tir
1308+
10181309
} // namespace tvm

0 commit comments

Comments
 (0)