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

tqchen 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 ae93d221d8 [RUNTIME][ROCM] Properly align rocm parameter buffer 
(#16550)
ae93d221d8 is described below

commit ae93d221d833bdc86df030ae53a12b21f7396fff
Author: Ruihang Lai <[email protected]>
AuthorDate: Sun Feb 11 10:03:08 2024 -0500

    [RUNTIME][ROCM] Properly align rocm parameter buffer (#16550)
    
    This PR properly aligns rocm param buffer fields to avoid bug in cases like
    
    struct Param {
      void *arg0;
      int32_t arg1;
      int64_t arg2;
    };
    
    In this case, we need to align arg2 by padding a i32 after arg1.
    
    Co-authored-by: Tianqi Chen <[email protected]>
---
 src/runtime/pack_args.h         | 27 ++++++++++++++++++++++-----
 src/runtime/rocm/rocm_module.cc |  2 +-
 2 files changed, 23 insertions(+), 6 deletions(-)

diff --git a/src/runtime/pack_args.h b/src/runtime/pack_args.h
index 5291b12fa1..9e8860d6ba 100644
--- a/src/runtime/pack_args.h
+++ b/src/runtime/pack_args.h
@@ -84,6 +84,9 @@ inline PackedFunc PackFuncNonBufferArg(F f, const 
std::vector<DLDataType>& arg_t
 /*!
  * \brief Create a packed function that from function that takes a packed 
arguments.
  *
+ * This procedure ensures inserts padding to ensure proper alignment of struct 
fields
+ * per C struct convention
+ *
  * \param f with signature (TVMArgs args, TVMRetValue* rv, void* pack_args, 
size_t nbytes)
  * \param arg_types The arguments that wish to get from
  * \tparam F the function type
@@ -91,7 +94,7 @@ inline PackedFunc PackFuncNonBufferArg(F f, const 
std::vector<DLDataType>& arg_t
  * \return The wrapped packed function.
  */
 template <typename F>
-inline PackedFunc PackFuncPackedArg(F f, const std::vector<DLDataType>& 
arg_types);
+inline PackedFunc PackFuncPackedArgAligned(F f, const std::vector<DLDataType>& 
arg_types);
 /*!
  * \brief Extract number of buffer argument from the argument types.
  * \param arg_types The argument types.
@@ -224,7 +227,7 @@ inline PackedFunc PackFuncNonBufferArg_(F f, int base, 
const std::vector<ArgConv
 }
 
 template <int N, typename F>
-inline PackedFunc PackFuncPackedArg_(F f, const std::vector<ArgConvertCode>& 
codes) {
+inline PackedFunc PackFuncPackedArgAligned_(F f, const 
std::vector<ArgConvertCode>& codes) {
   int num_args = static_cast<int>(codes.size());
   auto ret = [f, codes, num_args](TVMArgs args, TVMRetValue* ret) {
     TempArray<uint64_t, N> pack_(num_args);
@@ -232,30 +235,44 @@ inline PackedFunc PackFuncPackedArg_(F f, const 
std::vector<ArgConvertCode>& cod
     int32_t* ptr = pack;
     static_assert(sizeof(TVMValue) == 8, "invariant");
     static_assert(sizeof(void*) % sizeof(int32_t) == 0, "invariant");
+
+    // function to ensure alignment so fields are properly aligned
+    // factor: how many multiple of i32 we need to align to
+    auto ensure_alignment_to_multiple_of_i32 = [&](size_t factor) {
+      while ((ptr - pack) % factor != 0) {
+        ++ptr;
+      }
+    };
+
     for (int i = 0; i < num_args; ++i) {
       switch (codes[i]) {
         case HANDLE_TO_HANDLE: {
+          ensure_alignment_to_multiple_of_i32(sizeof(void*) / sizeof(int32_t));
           std::memcpy(ptr, &(args.values[i].v_handle), sizeof(void*));
           ptr += sizeof(void*) / sizeof(int32_t);
           break;
         }
         case INT64_TO_INT64:
         case FLOAT64_TO_FLOAT64: {
+          ensure_alignment_to_multiple_of_i32(2);
           std::memcpy(ptr, &args.values[i], sizeof(TVMValue));
           ptr += 2;
           break;
         }
         case INT64_TO_INT32: {
+          ensure_alignment_to_multiple_of_i32(1);
           *ptr = static_cast<int32_t>(args.values[i].v_int64);
           ++ptr;
           break;
         }
         case INT64_TO_UINT32: {
+          ensure_alignment_to_multiple_of_i32(1);
           *reinterpret_cast<uint32_t*>(ptr) = 
static_cast<uint32_t>(args.values[i].v_int64);
           ++ptr;
           break;
         }
         case FLOAT64_TO_FLOAT32: {
+          ensure_alignment_to_multiple_of_i32(1);
           *reinterpret_cast<float*>(ptr) = 
static_cast<float>(args.values[i].v_float64);
           ++ptr;
           break;
@@ -321,7 +338,7 @@ inline PackedFunc PackFuncNonBufferArg(F f, const 
std::vector<DLDataType>& arg_t
 }
 
 template <typename F>
-inline PackedFunc PackFuncPackedArg(F f, const std::vector<DLDataType>& 
arg_types) {
+inline PackedFunc PackFuncPackedArgAligned(F f, const std::vector<DLDataType>& 
arg_types) {
   std::vector<detail::ArgConvertCode> codes;
   for (size_t i = 0; i < arg_types.size(); ++i) {
     codes.push_back(detail::GetArgConvertCode(arg_types[i]));
@@ -329,9 +346,9 @@ inline PackedFunc PackFuncPackedArg(F f, const 
std::vector<DLDataType>& arg_type
   size_t nargs = codes.size();
   // specialization
   if (nargs <= 4) {
-    return detail::PackFuncPackedArg_<4>(f, codes);
+    return detail::PackFuncPackedArgAligned_<4>(f, codes);
   } else {
-    return detail::PackFuncPackedArg_<0>(f, codes);
+    return detail::PackFuncPackedArgAligned_<0>(f, codes);
   }
 }
 }  // namespace runtime
diff --git a/src/runtime/rocm/rocm_module.cc b/src/runtime/rocm/rocm_module.cc
index 9acd1ca903..96b5caa186 100644
--- a/src/runtime/rocm/rocm_module.cc
+++ b/src/runtime/rocm/rocm_module.cc
@@ -198,7 +198,7 @@ PackedFunc ROCMModuleNode::GetFunction(const String& name, 
const ObjectPtr<Objec
   const FunctionInfo& info = it->second;
   ROCMWrappedFunc f;
   f.Init(this, sptr_to_self, name, info.arg_types.size(), 
info.launch_param_tags);
-  return PackFuncPackedArg(f, info.arg_types);
+  return PackFuncPackedArgAligned(f, info.arg_types);
 }
 
 Module ROCMModuleCreate(std::string data, std::string fmt,

Reply via email to