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

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


The following commit(s) were added to refs/heads/test_all_cases_on_unity by 
this push:
     new b440f1b322 upd
b440f1b322 is described below

commit b440f1b3227f18fe8b74b64e9339765d810c2fe4
Author: Siyuan Feng <[email protected]>
AuthorDate: Tue Dec 5 10:54:25 2023 +0800

    upd
---
 include/tvm/topi/transform.h | 12 ++++++++----
 1 file changed, 8 insertions(+), 4 deletions(-)

diff --git a/include/tvm/topi/transform.h b/include/tvm/topi/transform.h
index 7ee09d834f..68f5d36c82 100644
--- a/include/tvm/topi/transform.h
+++ b/include/tvm/topi/transform.h
@@ -1591,16 +1591,20 @@ inline Tensor tensordot(const Tensor& A, const 
tvm::te::Tensor& B, Array<PrimExp
 
 inline Tensor arange(const PrimExpr& start, const PrimExpr& stop, const 
PrimExpr& step,
                      DataType dtype, std::string name = "T_arange", 
std::string tag = kInjective) {
+  arith::Analyzer analyzer;
   PrimExpr num_elem;
-  if (start.dtype().is_int() && stop.dtype().is_int() && 
step.dtype().is_int()) {
-    // fast path for integer arange
+  bool is_all_int = start.dtype().is_int() && stop.dtype().is_int() && 
step.dtype().is_int();
+  if (is_all_int && analyzer.CanProveGreaterEqual(step, 1)) {
+    // fast path for integer arange when step is positive
     num_elem = tvm::floordiv((stop - start + step - 1), step);
+  } else if (is_all_int && analyzer.CanProveLess(step, 0)) {
+    // fast path for integer arange when step is negative
+    num_elem = tvm::floordiv((start - stop - step - 1), -step);
   } else {
+    // fallback path for non-integer or step of unknown sign
     num_elem = tvm::cast(DefaultIndexType(),
                          tvm::ceil(tvm::cast(tvm::DataType::Float(32), stop - 
start) / step));
   }
-
-  arith::Analyzer analyzer;
   num_elem = analyzer.Simplify(num_elem);
 
   return compute(

Reply via email to