@@ -546,6 +546,7 @@ TileLayout::SplitMap TileLayout::GetSplitMap(
546546}
547547
548548/* ******* Normalization ********/
549+
549550using NodeSet = std::unordered_set<IterTreeBase, ObjectPtrHash, ObjectPtrEqual>;
550551
551552class 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+
7861075TileLayout 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