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 {