This is an automated email from the ASF dual-hosted git repository.

tqchen pushed a commit to branch unity
in repository https://gitbox.apache.org/repos/asf/tvm.git


The following commit(s) were added to refs/heads/unity by this push:
     new af14fbbbe1 [Relax] Fix to enable emit_te of topi scan/sort kernels 
(#16226)
af14fbbbe1 is described below

commit af14fbbbe1521170d8f6893aff1ffb15ee591f4e
Author: Bohan Hou <[email protected]>
AuthorDate: Tue Dec 12 09:19:21 2023 -0500

    [Relax] Fix to enable emit_te of topi scan/sort kernels (#16226)
    
    Fix to enable emit_te of topi scan/sort kernels
---
 python/tvm/dlight/gpu/fallback.py           | 10 ++++++++--
 src/arith/ir_visitor_with_analyzer.cc       |  2 +-
 src/tir/ir/data_type_rewriter.cc            |  2 +-
 src/tir/transforms/compact_buffer_region.cc |  2 +-
 src/tir/transforms/unify_thread_binding.cc  |  3 ++-
 5 files changed, 13 insertions(+), 6 deletions(-)

diff --git a/python/tvm/dlight/gpu/fallback.py 
b/python/tvm/dlight/gpu/fallback.py
index 3e0dbbcdaa..a800fc1acb 100644
--- a/python/tvm/dlight/gpu/fallback.py
+++ b/python/tvm/dlight/gpu/fallback.py
@@ -54,8 +54,14 @@ class Fallback(ScheduleRule):
             dom_kind = block.dom_kind()
             block = block.block_rv
 
-            if any(
-                [sch.get(loop_rv).thread_binding is not None for loop_rv in 
sch.get_loops(block)]
+            if (
+                any(
+                    [
+                        sch.get(loop_rv).thread_binding is not None
+                        for loop_rv in sch.get_loops(block)
+                    ]
+                )
+                or len(sch.get_loops(block)) == 0
             ):
                 continue
 
diff --git a/src/arith/ir_visitor_with_analyzer.cc 
b/src/arith/ir_visitor_with_analyzer.cc
index e7cf3ea7ea..dba4567f88 100644
--- a/src/arith/ir_visitor_with_analyzer.cc
+++ b/src/arith/ir_visitor_with_analyzer.cc
@@ -68,7 +68,7 @@ void IRVisitorWithAnalyzer::VisitStmt_(const AttrStmtNode* 
op) {
   if (op->attr_key == tir::attr::thread_extent || op->attr_key == 
tir::attr::virtual_thread) {
     IterVar iv = Downcast<IterVar>(op->node);
     ICHECK_NE(iv->thread_tag.length(), 0U);
-    analyzer_.Bind(iv->var, Range::FromMinExtent(0, op->value));
+    analyzer_.Bind(iv->var, Range::FromMinExtent(IntImm(op->value->dtype, 0), 
op->value));
   }
   StmtExprVisitor::VisitStmt_(op);
 }
diff --git a/src/tir/ir/data_type_rewriter.cc b/src/tir/ir/data_type_rewriter.cc
index 619804b880..aa8f2f3f6f 100644
--- a/src/tir/ir/data_type_rewriter.cc
+++ b/src/tir/ir/data_type_rewriter.cc
@@ -610,7 +610,7 @@ bool IndexDataTypeNormalizer::CanRewriteDType(DataType 
dtype) const {
 }
 
 PrimExpr IndexDataTypeNormalizer::VisitExpr_(const IntImmNode* op) {
-  if (is_enabled_) {
+  if (is_enabled_ && CanRewriteDType(op->dtype)) {
     ICHECK_LE(op->value, 
Downcast<Integer>(max_value(target_data_type_))->value);
     return cast(target_data_type_, GetRef<IntImm>(op));
   }
diff --git a/src/tir/transforms/compact_buffer_region.cc 
b/src/tir/transforms/compact_buffer_region.cc
index 6acfe02a10..c7706212c5 100644
--- a/src/tir/transforms/compact_buffer_region.cc
+++ b/src/tir/transforms/compact_buffer_region.cc
@@ -296,7 +296,7 @@ class BufferAccessRegionCollector : public StmtExprVisitor {
       ancestor_iters_.push_back(iter);
       Range dom = iter->dom;
       if (!dom.defined()) {  // dom is empty for legacy te schedule
-        dom = Range::FromMinExtent(0, op->value);
+        dom = Range::FromMinExtent(make_zero(op->value->dtype), op->value);
       }
       dom_analyzer_.Bind(iter->var, dom);
       dom_map_.emplace(iter->var.get(), arith::IntSet::FromRange(dom));
diff --git a/src/tir/transforms/unify_thread_binding.cc 
b/src/tir/transforms/unify_thread_binding.cc
index 09b0970dd3..02fa333dbe 100644
--- a/src/tir/transforms/unify_thread_binding.cc
+++ b/src/tir/transforms/unify_thread_binding.cc
@@ -50,7 +50,8 @@ class ThreadBindingUnifier : public StmtExprMutator {
       return StmtMutator::VisitStmt_(op);
     }
     IterVar old_iter_var = Downcast<IterVar>(op->node);
-    return UnifyThreadBindingImpl(op, old_iter_var->var, old_iter_var, 
old_iter_var->dom);
+    return UnifyThreadBindingImpl(op, old_iter_var->var, old_iter_var,
+                                  
Range::FromMinExtent(IntImm(op->value->dtype, 0), op->value));
   }
 
   Stmt VisitStmt_(const ForNode* op) final {

Reply via email to