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

masahi 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 d3d350b858 [Unity] cuda graph support for cublas (#15435)
d3d350b858 is described below

commit d3d350b85894e898c2376d40d5c9afd83ba5b6e8
Author: Sunghyun Park <[email protected]>
AuthorDate: Sat Jul 29 18:41:10 2023 -0700

    [Unity] cuda graph support for cublas (#15435)
    
    * cuda graph support for cublas
    
    * reflect feedback
---
 src/runtime/contrib/cublas/cublas.cc              | 16 ++++---
 src/runtime/contrib/cublas/cublas_json_runtime.cc |  8 +++-
 src/runtime/contrib/cublas/cublas_utils.h         |  4 +-
 tests/python/relax/test_codegen_cublas.py         | 51 ++++++++++++++++++++---
 4 files changed, 64 insertions(+), 15 deletions(-)

diff --git a/src/runtime/contrib/cublas/cublas.cc 
b/src/runtime/contrib/cublas/cublas.cc
index b49f15008c..7a8637b329 100644
--- a/src/runtime/contrib/cublas/cublas.cc
+++ b/src/runtime/contrib/cublas/cublas.cc
@@ -135,8 +135,9 @@ int roundoff(int v, int d) { return (v + d - 1) / d * d; }
 
 #if CUDART_VERSION >= 10010
 
-void CallCublasLt(cublasLtHandle_t hdl, const DLTensor* A, const DLTensor* B, 
const DLTensor* bias,
-                  const DLTensor* C, bool transa, bool transb, 
cublasLtEpilogue_t epilogue) {
+void CallCublasLt(cublasLtHandle_t hdl, cudaStream_t stream, const DLTensor* 
A, const DLTensor* B,
+                  const DLTensor* bias, const DLTensor* C, bool transa, bool 
transb,
+                  cublasLtEpilogue_t epilogue) {
   ICHECK(TypeEqual(A->dtype, B->dtype));
   // Reversed strides indicates an in-place transpose operation.
   transa = IsInPlaceTransposed(A) ? !transa : transa;
@@ -240,7 +241,7 @@ void CallCublasLt(cublasLtHandle_t hdl, const DLTensor* A, 
const DLTensor* B, co
   auto C_data = static_cast<char*>(C->data) + C->byte_offset;
 
   CHECK_CUBLAS_ERROR(cublasLtMatmul(hdl, op_desc, alpha, B_data, A_desc, 
A_data, B_desc, beta,
-                                    C_data, C_desc, C_data, C_desc, nullptr, 
nullptr, 0, nullptr));
+                                    C_data, C_desc, C_data, C_desc, nullptr, 
nullptr, 0, stream));
 
   cublasLtMatmulDescDestroy(op_desc);
   cublasLtMatrixLayoutDestroy(A_desc);
@@ -248,7 +249,7 @@ void CallCublasLt(cublasLtHandle_t hdl, const DLTensor* A, 
const DLTensor* B, co
   cublasLtMatrixLayoutDestroy(C_desc);
 }
 
-inline void CallLtIgemm(TVMArgs args, TVMRetValue* ret, cublasLtHandle_t hdl) {
+inline void CallLtIgemm(TVMArgs args, TVMRetValue* ret, cublasLtHandle_t hdl, 
cudaStream_t stream) {
   DLTensor* A = args[0];
   DLTensor* B = args[1];
   DLTensor* C = args[2];
@@ -316,7 +317,7 @@ inline void CallLtIgemm(TVMArgs args, TVMRetValue* ret, 
cublasLtHandle_t hdl) {
                                                       &order_COL32, 
sizeof(order_COL32)));
 
   CHECK_CUBLAS_ERROR(cublasLtMatmul(hdl, operationDesc, &alpha, B_data, Adesc, 
A_data, Bdesc, &beta,
-                                    C_data, Cdesc, C_data, Cdesc, nullptr, 
nullptr, 0, nullptr));
+                                    C_data, Cdesc, C_data, Cdesc, nullptr, 
nullptr, 0, stream));
 }
 #endif
 
@@ -490,7 +491,10 @@ 
TVM_REGISTER_GLOBAL("tvm.contrib.cublaslt.matmul").set_body([](TVMArgs args, TVM
   ICHECK(TypeMatch(A->dtype, kDLInt, 8)) << "Expects dtype to be int8\n";
   cublasLtHandle_t ltHandle;
   CHECK_CUBLAS_ERROR(cublasLtCreate(&ltHandle));
-  CallLtIgemm(args, ret, ltHandle);
+  auto func = tvm::runtime::Registry::Get("runtime.get_cuda_stream");
+  ICHECK(func != nullptr);
+  cudaStream_t stream = static_cast<cudaStream_t>((*func)().operator void*());
+  CallLtIgemm(args, ret, ltHandle, stream);
   CHECK_CUBLAS_ERROR(cublasLtDestroy(ltHandle));
 });
 #endif  // CUDART_VERSION >= 10010
diff --git a/src/runtime/contrib/cublas/cublas_json_runtime.cc 
b/src/runtime/contrib/cublas/cublas_json_runtime.cc
index fc6b2f5e62..9617559d7e 100644
--- a/src/runtime/contrib/cublas/cublas_json_runtime.cc
+++ b/src/runtime/contrib/cublas/cublas_json_runtime.cc
@@ -54,6 +54,10 @@ class CublasJSONRuntime : public JSONRuntimeBase {
   void Run() override {
     auto* entry_ptr = tvm::contrib::CuBlasLtThreadEntry::ThreadLocal();
 
+    auto func = tvm::runtime::Registry::Get("runtime.get_cuda_stream");
+    ICHECK(func != nullptr);
+    cudaStream_t stream = static_cast<cudaStream_t>((*func)().operator 
void*());
+
     for (size_t i = 0; i < nodes_.size(); ++i) {
       const auto& node = nodes_[i];
       if (node.GetOpType() == "kernel") {
@@ -86,8 +90,8 @@ class CublasJSONRuntime : public JSONRuntimeBase {
 
         auto [a_ptr, b_ptr, bias_ptr] = get_inputs(node, epilogue != 
CUBLASLT_EPILOGUE_DEFAULT);
 
-        tvm::contrib::CallCublasLt(entry_ptr->handle, a_ptr, b_ptr, bias_ptr, 
out_ptr, transa,
-                                   transb, epilogue);
+        tvm::contrib::CallCublasLt(entry_ptr->handle, stream, a_ptr, b_ptr, 
bias_ptr, out_ptr,
+                                   transa, transb, epilogue);
       }
     }
   }
diff --git a/src/runtime/contrib/cublas/cublas_utils.h 
b/src/runtime/contrib/cublas/cublas_utils.h
index 18ea8de8ef..82c77cfb5e 100644
--- a/src/runtime/contrib/cublas/cublas_utils.h
+++ b/src/runtime/contrib/cublas/cublas_utils.h
@@ -113,8 +113,8 @@ inline cudaDataType_t GetCudaDataType(DLDataType type) {
 }
 
 /*! \brief Execute matrix multiply followed by the specified epilogue, using 
cuBLASLt. */
-void CallCublasLt(cublasLtHandle_t hdl, const DLTensor* A, const DLTensor* B, 
const DLTensor* bias,
-                  const DLTensor* C, bool transa, bool transb,
+void CallCublasLt(cublasLtHandle_t hdl, cudaStream_t stream, const DLTensor* 
A, const DLTensor* B,
+                  const DLTensor* bias, const DLTensor* C, bool transa, bool 
transb,
                   cublasLtEpilogue_t epilogue = CUBLASLT_EPILOGUE_DEFAULT);
 
 }  // namespace contrib
diff --git a/tests/python/relax/test_codegen_cublas.py 
b/tests/python/relax/test_codegen_cublas.py
index 4eb0cc3b0a..68e625a232 100644
--- a/tests/python/relax/test_codegen_cublas.py
+++ b/tests/python/relax/test_codegen_cublas.py
@@ -41,23 +41,30 @@ cublas_enabled = pytest.mark.skipif(
 pytestmark = [cublas_enabled]
 
 
-def build_and_run(mod, inputs_np, target, legalize=False):
+def build_and_run(mod, inputs_np, target, legalize=False, cuda_graph=False):
     if legalize:
         mod = relax.transform.LegalizeOps()(mod)
 
     dev = tvm.device(target, 0)
-    ex = relax.build(mod, target)
+    with tvm.transform.PassContext(config={"relax.backend.use_cuda_graph": 
cuda_graph}):
+        ex = relax.build(mod, target)
     vm = relax.VirtualMachine(ex, dev)
     f = vm["main"]
     inputs = [tvm.nd.array(inp, dev) for inp in inputs_np]
+
+    # For cuda graph, run the compiled function twice to make sure that we can 
launch the cached
+    # graph on the second run.
+    if cuda_graph:
+        f(*inputs)
+
     return f(*inputs).numpy()
 
 
-def get_result_with_relax_cublas_offload(mod, *args):
+def get_result_with_relax_cublas_offload(mod, np_inputs, cuda_graph=False):
     mod = partition_for_cublas(mod)
     mod = relax.transform.RunCodegen()(mod)
 
-    return build_and_run(mod, args, "cuda")
+    return build_and_run(mod, np_inputs, "cuda", cuda_graph)
 
 
 def _to_concrete_shape(symbolic_shape, var_table):
@@ -146,7 +153,7 @@ def test_matmul_offload(
         activation=activation,
     )
 
-    out = get_result_with_relax_cublas_offload(mod, *args)
+    out = get_result_with_relax_cublas_offload(mod, args)
     ref = build_and_run(mod, args, "llvm", legalize=True)
 
     tvm.testing.assert_allclose(out, ref, rtol=1e-2, atol=1e-2)
@@ -161,5 +168,39 @@ def test_cublass_partition_matmul_without_bias():
     assert len(mod["main"].body.blocks[0].bindings) == 2
 
 
+def test_cublas_matmul_cuda_graph():
+    @tvm.script.ir.ir_module
+    class Mod:
+        @R.function
+        def main(
+            x: R.Tensor((16, 16), "float16"),
+            w0: R.Tensor((16, 16), "float16"),
+            w1: R.Tensor((16, 16), "float16"),
+            w2: R.Tensor((16, 16), "float16"),
+        ):
+            R.func_attr({"num_input": 1})
+            with R.dataflow():
+                lv0 = R.matmul(x, w0)
+                lv1 = R.matmul(lv0, w1)
+                lv2 = R.matmul(lv1, w2)
+                R.output(lv2)
+            return lv2
+
+    mod = Mod
+    shape = [16, 16]
+    data = np.random.rand(*shape).astype(np.float16)
+    w0 = np.random.rand(*shape).astype(np.float16)
+    w1 = np.random.rand(*shape).astype(np.float16)
+    w2 = np.random.rand(*shape).astype(np.float16)
+    inputs = (data, w0, w1, w2)
+
+    out = get_result_with_relax_cublas_offload(Mod, inputs, cuda_graph=True)
+
+    with tvm.target.Target("cuda"):
+        mod = tvm.tir.transform.DefaultGPUSchedule()(mod)
+    ref = build_and_run(mod, inputs, "llvm", legalize=True)
+    tvm.testing.assert_allclose(out, ref, rtol=1e-2, atol=1e-2)
+
+
 if __name__ == "__main__":
     tvm.testing.main()

Reply via email to