Use rocprofiler-sdk for RCCL-API tracing (#126)

- Add support for RCCL API tracing through rocprofiler-sdk.
- Refactored the comm_data code to use the SDK RCCL_API callbacks.
- Add a runtime version check for SDK to gate callback enablement, rather than just the compile-time check.
- Fixed: SAMPLING_TIMEOUT was not being handled correctly in add_test.
This commit is contained in:
David Galiffi
2025-06-06 11:36:17 -04:00
کامیت شده توسط GitHub
والد c5507e3740
کامیت af77d93f75
31فایلهای تغییر یافته به همراه350 افزوده شده و 1252 حذف شده
@@ -34,6 +34,7 @@
#include "library/components/category_region.hpp"
#include "library/rocprofiler-sdk/counters.hpp"
#include "library/rocprofiler-sdk/fwd.hpp"
#include "library/rocprofiler-sdk/rccl.hpp"
#include "library/thread_info.hpp"
#include "library/tracing.hpp"
@@ -482,6 +483,17 @@ tool_tracing_callback(rocprofiler_callback_tracing_record_t record,
auto ts = rocprofiler_timestamp_t{};
ROCPROFILER_CALL(rocprofiler_get_timestamp(&ts));
const char* name = nullptr;
rocprofiler_query_callback_tracing_kind_operation_name(record.kind, record.operation,
&name, nullptr);
auto info = std::stringstream{};
info << std::left << "tid=" << record.thread_id << ", cid=" << std::setw(3)
<< record.correlation_id.internal << ", kind=" << std::setw(2) << record.kind
<< ", operation=" << std::setw(3) << record.operation
<< ", phase=" << record.phase << ", dt_nsec=" << std::setw(8) << ts
<< ", name=" << name;
if(record.phase == ROCPROFILER_CALLBACK_PHASE_ENTER)
{
user_data->value = ts;
@@ -525,6 +537,12 @@ tool_tracing_callback(rocprofiler_callback_tracing_record_t record,
break;
}
#endif
case ROCPROFILER_CALLBACK_TRACING_RCCL_API:
{
tool_tracing_callback_start(category::rocm_rccl_api{}, record, user_data,
ts);
break;
}
case ROCPROFILER_CALLBACK_TRACING_NONE:
case ROCPROFILER_CALLBACK_TRACING_LAST:
case ROCPROFILER_CALLBACK_TRACING_MARKER_CONTROL_API:
@@ -533,7 +551,6 @@ tool_tracing_callback(rocprofiler_callback_tracing_record_t record,
case ROCPROFILER_CALLBACK_TRACING_SCRATCH_MEMORY:
case ROCPROFILER_CALLBACK_TRACING_KERNEL_DISPATCH:
case ROCPROFILER_CALLBACK_TRACING_MEMORY_COPY:
case ROCPROFILER_CALLBACK_TRACING_RCCL_API:
#if(ROCPROFILER_VERSION >= 600)
case ROCPROFILER_CALLBACK_TRACING_OMPT:
case ROCPROFILER_CALLBACK_TRACING_MEMORY_ALLOCATION:
@@ -616,6 +633,13 @@ tool_tracing_callback(rocprofiler_callback_tracing_record_t record,
break;
}
#endif
case ROCPROFILER_CALLBACK_TRACING_RCCL_API:
{
tool_tracing_callback_rccl(record, user_data->value, ts);
tool_tracing_callback_stop(category::rocm_rccl_api{}, record, user_data,
ts, _bt_data);
break;
}
case ROCPROFILER_CALLBACK_TRACING_NONE:
case ROCPROFILER_CALLBACK_TRACING_LAST:
case ROCPROFILER_CALLBACK_TRACING_MARKER_CONTROL_API:
@@ -624,7 +648,6 @@ tool_tracing_callback(rocprofiler_callback_tracing_record_t record,
case ROCPROFILER_CALLBACK_TRACING_SCRATCH_MEMORY:
case ROCPROFILER_CALLBACK_TRACING_KERNEL_DISPATCH:
case ROCPROFILER_CALLBACK_TRACING_MEMORY_COPY:
case ROCPROFILER_CALLBACK_TRACING_RCCL_API:
#if(ROCPROFILER_VERSION >= 600)
case ROCPROFILER_CALLBACK_TRACING_OMPT:
case ROCPROFILER_CALLBACK_TRACING_MEMORY_ALLOCATION:
@@ -651,6 +674,12 @@ tool_tracing_callback(rocprofiler_callback_tracing_record_t record,
_data->dispatch_info.dispatch_id,
timing_interval{ _data->start_timestamp, _data->end_timestamp });
}
else
{
ROCPROFSYS_WARNING_F(
1, "tool_tracing_callback: unhandled PHASE_NONE callback record\n\t%s\n",
info.str().c_str());
}
}
else
{
@@ -1025,13 +1054,16 @@ tool_init(rocprofiler_client_finalize_t fini_func, void* user_data)
{
auto domains = settings::instance()->at("ROCPROFSYS_ROCM_DOMAINS");
ROCPROFSYS_VERBOSE_F(1, "rocprof-sys ROCm Domains:\n");
ROCPROFSYS_VERBOSE_F(1, "Available ROCm Domains:\n");
for(const auto& itr : domains->get_choices())
ROCPROFSYS_VERBOSE_F(1, "- %s\n", itr.c_str());
auto _callback_domains = rocprofiler_sdk::get_callback_domains();
auto _buffered_domain = rocprofiler_sdk::get_buffered_domains();
auto _counter_events = rocprofiler_sdk::get_rocm_events();
auto _version = rocprofiler_sdk::get_version();
ROCPROFSYS_WARNING_IF(_version.formatted == 0,
"Warning! rocprofiler-sdk version not initialized\n");
auto* _data = as_client_data(user_data);
_data->client_fini = fini_func;
@@ -1052,11 +1084,14 @@ tool_init(rocprofiler_client_finalize_t fini_func, void* user_data)
ROCPROFILER_CALLBACK_TRACING_HSA_FINALIZE_EXT_API,
ROCPROFILER_CALLBACK_TRACING_HIP_RUNTIME_API,
ROCPROFILER_CALLBACK_TRACING_HIP_COMPILER_API,
#if(ROCPROFILER_VERSION >= 700)
ROCPROFILER_CALLBACK_TRACING_MARKER_CORE_API,
ROCPROFILER_CALLBACK_TRACING_RCCL_API,
#if(ROCPROFILER_VERSION >= 600)
ROCPROFILER_CALLBACK_TRACING_ROCDECODE_API,
#endif
#if(ROCPROFILER_VERSION >= 700)
ROCPROFILER_CALLBACK_TRACING_ROCJPEG_API,
#endif
ROCPROFILER_CALLBACK_TRACING_MARKER_CORE_API
})
{
if(_callback_domains.count(itr) > 0)