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

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


The following commit(s) were added to refs/heads/main by this push:
     new 5308739741 [TIR] Allow sync threads inside condition (#16345)
5308739741 is described below

commit 5308739741bc9962ecbbcd7d58182f0874508c19
Author: Bohan Hou <[email protected]>
AuthorDate: Thu Jan 4 12:53:25 2024 -0500

    [TIR] Allow sync threads inside condition (#16345)
    
    Originally, it is not allowed to sync threads inside a condition `while, 
if`.
    
    This PR introduces `tvm_thread_invariant` op to annotate the condition to 
be thread id invariant and get around the check.
---
 include/tvm/tir/builtin.h                        |  6 ++++
 python/tvm/script/ir_builder/tir/ir.py           |  2 ++
 python/tvm/tir/op.py                             | 17 +++++++++
 src/target/source/codegen_c.cc                   |  4 +++
 src/tir/op/builtin.cc                            |  4 +++
 src/tir/transforms/storage_access.cc             | 30 +++++++++++++---
 tests/python/codegen/test_target_codegen_cuda.py | 46 ++++++++++++++++++++++++
 7 files changed, 105 insertions(+), 4 deletions(-)

diff --git a/include/tvm/tir/builtin.h b/include/tvm/tir/builtin.h
index 65012c6c0f..96222e03a9 100644
--- a/include/tvm/tir/builtin.h
+++ b/include/tvm/tir/builtin.h
@@ -411,6 +411,12 @@ TVM_DLL const Op& tvm_check_return();
  */
 TVM_DLL const Op& tvm_thread_context();
 
+/*!
+ * \brief Mark a condition to be thread invariant.
+ *  This means the condition must be the same for all threads.
+ */
+TVM_DLL const Op& tvm_thread_invariant();
+
 /*!
  * \brief Lowered version of call packed, the space of value and
  *  type codes are explicitly allocated.
diff --git a/python/tvm/script/ir_builder/tir/ir.py 
b/python/tvm/script/ir_builder/tir/ir.py
index d4a7445b7d..b5f427c34c 100644
--- a/python/tvm/script/ir_builder/tir/ir.py
+++ b/python/tvm/script/ir_builder/tir/ir.py
@@ -1832,6 +1832,7 @@ call_cpacked_lowered = 
_op_wrapper(_tir_op.call_cpacked_lowered)
 tvm_tuple = _op_wrapper(_tir_op.tvm_tuple)
 tvm_struct_set = _op_wrapper(_tir_op.tvm_struct_set)
 tvm_struct_get = _tir_op.tvm_struct_get
+tvm_thread_invariant = _op_wrapper(_tir_op.tvm_thread_invariant)
 tvm_thread_allreduce = _op_wrapper(_tir_op.tvm_thread_allreduce)
 tvm_load_matrix_sync = _op_wrapper(_tir_op.tvm_load_matrix_sync)
 tvm_mma_sync = _op_wrapper(_tir_op.tvm_mma_sync)
@@ -2104,6 +2105,7 @@ __all__ = [
     "tvm_tuple",
     "tvm_struct_set",
     "tvm_struct_get",
+    "tvm_thread_invariant",
     "tvm_thread_allreduce",
     "tvm_load_matrix_sync",
     "tvm_mma_sync",
diff --git a/python/tvm/tir/op.py b/python/tvm/tir/op.py
index bb2530b125..d7478645c5 100644
--- a/python/tvm/tir/op.py
+++ b/python/tvm/tir/op.py
@@ -602,6 +602,23 @@ def tvm_thread_allreduce(*freduce_args):
     return call_intrin("handle", "tir.tvm_thread_allreduce", *freduce_args)
 
 
+def tvm_thread_invariant(cond):
+    """Mark condition as thread invariant.
+
+    Parameters
+    ----------
+    cond : Expr
+        The condition.
+
+    Returns
+    -------
+    call : PrimExpr
+        The call expression.
+    """
+    assert isinstance(cond, PrimExpr)
+    return call_intrin(cond.dtype, "tir.tvm_thread_invariant", cond)
+
+
 def tvm_storage_sync(storage_scope):
     """Perform synchronization in specified scope.
 
diff --git a/src/target/source/codegen_c.cc b/src/target/source/codegen_c.cc
index 0ff0531b5c..8380971249 100644
--- a/src/target/source/codegen_c.cc
+++ b/src/target/source/codegen_c.cc
@@ -669,6 +669,10 @@ void CodeGenC::VisitExpr_(const CallNode* op, 
std::ostream& os) {  // NOLINT(*)
       const StringImmNode* str = op->args[0].as<StringImmNode>();
       ICHECK(str != nullptr);
       os << "__tvm_param__" << str->value;
+    } else if (op->op.same_as(builtin::tvm_thread_invariant())) {
+      os << "(";
+      this->PrintExpr(op->args[0], os);
+      os << ")";
     } else {
       LOG(FATAL) << "Unresolved call " << op->op;
     }
diff --git a/src/tir/op/builtin.cc b/src/tir/op/builtin.cc
index 1b80959b57..a5089e2566 100644
--- a/src/tir/op/builtin.cc
+++ b/src/tir/op/builtin.cc
@@ -211,6 +211,10 @@ TIR_DEFINE_BUILTIN_FUNC(tvm_thread_context)
     .set_num_inputs(1)
     .set_attr<TCallEffectKind>("TCallEffectKind", 
Integer(CallEffectKind::kOpaque));
 
+TIR_DEFINE_BUILTIN_FUNC(tvm_thread_invariant)
+    .set_num_inputs(1)
+    .set_attr<TCallEffectKind>("TCallEffectKind", 
Integer(CallEffectKind::kPure));
+
 TIR_DEFINE_BUILTIN_FUNC(tvm_call_packed_lowered)
     .set_attr<TCallEffectKind>("TCallEffectKind", 
Integer(CallEffectKind::kOpaque))
     .set_attr<TScriptPrinterName>("TScriptPrinterName", 
String("call_packed_lowered"),
diff --git a/src/tir/transforms/storage_access.cc 
b/src/tir/transforms/storage_access.cc
index b34cfdfb31..cbc7f07cae 100644
--- a/src/tir/transforms/storage_access.cc
+++ b/src/tir/transforms/storage_access.cc
@@ -170,8 +170,23 @@ void StorageAccessVisitor::VisitStmt_(const ForNode* op) {
   }
 }
 
+bool IsThreadInvariant(const PrimExpr& cond) {
+  if (auto call = cond.as<CallNode>()) {
+    if (auto opt_call_op = call->op.as<Op>()) {
+      auto call_op = opt_call_op.value();
+      if (call_op.same_as(builtin::tvm_thread_invariant())) {
+        return true;
+      }
+    }
+  }
+  return false;
+}
+
 void StorageAccessVisitor::VisitStmt_(const IfThenElseNode* op) {
-  ++condition_counter_;
+  bool is_thread_invariant = IsThreadInvariant(op->condition);
+  if (!is_thread_invariant) {
+    ++condition_counter_;
+  }
   this->VisitExpr(op->condition);
   scope_.push_back(std::vector<StmtEntry>());
   this->VisitStmt(op->then_case);
@@ -187,11 +202,16 @@ void StorageAccessVisitor::VisitStmt_(const 
IfThenElseNode* op) {
     s.access.insert(s.access.end(), v.begin(), v.end());
   }
   scope_.back().emplace_back(std::move(s));
-  --condition_counter_;
+  if (!is_thread_invariant) {
+    --condition_counter_;
+  }
 }
 
 void StorageAccessVisitor::VisitStmt_(const WhileNode* op) {
-  ++condition_counter_;
+  bool is_thread_invariant = IsThreadInvariant(op->condition);
+  if (!is_thread_invariant) {
+    ++condition_counter_;
+  }
   this->VisitExpr(op->condition);
   scope_.push_back(std::vector<StmtEntry>());
   this->VisitStmt(op->body);
@@ -200,7 +220,9 @@ void StorageAccessVisitor::VisitStmt_(const WhileNode* op) {
   s.access = Summarize(std::move(scope_.back()), nullptr);
   scope_.pop_back();
   scope_.back().emplace_back(std::move(s));
-  --condition_counter_;
+  if (!is_thread_invariant) {
+    --condition_counter_;
+  }
 }
 
 void StorageAccessVisitor::VisitExpr_(const CallNode* op) {
diff --git a/tests/python/codegen/test_target_codegen_cuda.py 
b/tests/python/codegen/test_target_codegen_cuda.py
index 9b14af5abf..3b5cc00019 100644
--- a/tests/python/codegen/test_target_codegen_cuda.py
+++ b/tests/python/codegen/test_target_codegen_cuda.py
@@ -24,6 +24,7 @@ import numpy as np
 from tvm import topi
 from tvm.contrib.nvcc import have_fp16, have_int8, have_bf16
 from tvm.contrib import utils
+from tvm.script import tir as T
 import tvm.testing
 import pytest
 
@@ -1068,5 +1069,50 @@ def test_cuda_save_kernels_for_profiling():
     check_cuda(64, 2)
 
 
+def test_cuda_thread_sync_inside_condition():
+    @T.prim_func
+    def func1(A: T.Buffer((4, 4), "float32")) -> None:
+        A_shared = T.alloc_buffer((4, 4), "float32", scope="shared")
+        for bx in T.thread_binding(1, "blockIdx.x"):
+            for tx in T.thread_binding(32, "threadIdx.x"):
+                if A[0, 0] > 1.0:
+                    for i, j in T.grid(4, 4):
+                        A_shared[i, j] = A[i, j]
+                    for i, j in T.grid(4, 4):
+                        A[i, j] = A_shared[i, j] + 1.0
+
+    @T.prim_func
+    def func2(A: T.Buffer((4, 4), "float32")) -> None:
+        A_shared = T.alloc_buffer((4, 4), "float32", scope="shared")
+        for bx in T.thread_binding(1, "blockIdx.x"):
+            for tx in T.thread_binding(32, "threadIdx.x"):
+                if T.tvm_thread_invariant(A[0, 0] > 1.0):
+                    for i, j in T.grid(4, 4):
+                        A_shared[i, j] = A[i, j]
+                    for i, j in T.grid(4, 4):
+                        A[i, j] = A_shared[i, j] + 1.0
+
+    @T.prim_func
+    def func3(A: T.Buffer((4, 4), "float32")) -> None:
+        A_shared = T.alloc_buffer((4, 4), "float32", scope="shared")
+        for bx in T.thread_binding(1, "blockIdx.x"):
+            for tx in T.thread_binding(32, "threadIdx.x"):
+                while T.tvm_thread_invariant(A[0, 0] > 1.0):
+                    for i, j in T.grid(4, 4):
+                        A_shared[i, j] = A[i, j]
+                    for i, j in T.grid(4, 4):
+                        A[i, j] = A_shared[i, j] + 1.0
+
+    mod = tvm.IRModule({"main": func1})
+    with pytest.raises(tvm.error.InternalError):
+        tvm.build(mod, target="cuda")
+
+    mod = tvm.IRModule({"main": func2})
+    tvm.build(mod, target="cuda")
+
+    mod = tvm.IRModule({"main": func3})
+    tvm.build(mod, target="cuda")
+
+
 if __name__ == "__main__":
     tvm.testing.main()

Reply via email to