ATT Agent fixes and improvements (#1011)

* Tidying ATT dispatch API. ATT Agent to be initialized with rest of profiler. Removing read_index-based wait.

* Formatting

* Adding some input validation

* Add perf test for agent

* Removing async
This commit is contained in:
Giovanni Lenzi Baraldi
2024-08-15 13:57:13 -03:00
committed by GitHub
parent 5bec9f66b6
commit fa1b9e67ab
11 changed files with 286 additions and 139 deletions
@@ -40,8 +40,6 @@ ROCPROFILER_EXTERN_C_INIT
typedef enum
{
ROCPROFILER_ATT_CONTROL_NONE = 0,
ROCPROFILER_ATT_CONTROL_START = 1,
ROCPROFILER_ATT_CONTROL_STOP = 2,
ROCPROFILER_ATT_CONTROL_START_AND_STOP = 3
} rocprofiler_att_control_flags_t;
@@ -321,7 +321,8 @@ start_context(rocprofiler_context_id_t context_id)
auto status = ROCPROFILER_STATUS_SUCCESS;
if(cfg->counter_collection) rocprofiler::counters::start_context(cfg);
if(cfg->thread_trace) cfg->thread_trace->start_context();
if(cfg->agent_thread_trace) cfg->agent_thread_trace->start_context();
if(cfg->dispatch_thread_trace) cfg->dispatch_thread_trace->start_context();
if(cfg->agent_counter_collection) status = rocprofiler::counters::start_agent_ctx(cfg);
#if ROCPROFILER_SDK_HSA_PC_SAMPLING > 0
if(cfg->pc_sampler) status = rocprofiler::pc_sampling::start_service(cfg);
@@ -355,7 +356,9 @@ stop_context(rocprofiler_context_id_t idx)
rocprofiler::counters::stop_context(const_cast<context*>(_expected));
}
if(_expected->thread_trace) _expected->thread_trace->stop_context();
if(_expected->agent_thread_trace) _expected->agent_thread_trace->stop_context();
if(_expected->dispatch_thread_trace)
_expected->dispatch_thread_trace->stop_context();
if(_expected->agent_counter_collection)
{
@@ -128,7 +128,8 @@ struct context
std::unique_ptr<agent_counter_collection_service> agent_counter_collection = {};
std::unique_ptr<pc_sampling_service> pc_sampler = {};
std::unique_ptr<thread_trace::ThreadTracerInterface> thread_trace = {};
std::unique_ptr<thread_trace::DispatchThreadTracer> dispatch_thread_trace = {};
std::unique_ptr<thread_trace::AgentThreadTracer> agent_thread_trace = {};
};
// set the client index needs to be called before allocate_context()
+9 -10
View File
@@ -36,6 +36,9 @@ namespace rocprofiler
{
namespace hsa
{
constexpr uint16_t VENDOR_BIT = HSA_PACKET_TYPE_VENDOR_SPECIFIC << HSA_PACKET_HEADER_TYPE;
constexpr uint16_t BARRIER_BIT = 1 << HSA_PACKET_HEADER_BARRIER;
hsa_status_t
CounterAQLPacket::CounterMemoryPool::Alloc(void** ptr, size_t size, desc_t flags, void* data)
{
@@ -114,13 +117,10 @@ CounterAQLPacket::CounterAQLPacket(aqlprofile_agent_handle_t ag
reinterpret_cast<void*>(&pool));
if(status != HSA_STATUS_SUCCESS) ROCP_FATAL << "Could not create PMC packets!";
packets.start_packet.header = 1 << HSA_PACKET_HEADER_BARRIER;
// Read and stop packets require the barrier bit set to wait for the dispatch packet
auto header = (HSA_PACKET_TYPE_VENDOR_SPECIFIC << HSA_PACKET_HEADER_TYPE) |
(1 << HSA_PACKET_HEADER_BARRIER);
packets.stop_packet.header = header;
packets.read_packet.header = header;
empty = false;
packets.start_packet.header = VENDOR_BIT;
packets.stop_packet.header = VENDOR_BIT | BARRIER_BIT;
packets.read_packet.header = VENDOR_BIT | BARRIER_BIT;
empty = false;
}
hsa_status_t
@@ -181,9 +181,8 @@ TraceControlAQLPacket::TraceControlAQLPacket(const TraceMemoryPool& _tr
tracepool.get());
CHECK_HSA(status, "failed to create ATT packet");
packets.start_packet.header = HSA_PACKET_TYPE_VENDOR_SPECIFIC << HSA_PACKET_HEADER_TYPE;
packets.stop_packet.header = (HSA_PACKET_TYPE_VENDOR_SPECIFIC << HSA_PACKET_HEADER_TYPE) |
(1 << HSA_PACKET_HEADER_BARRIER);
packets.start_packet.header = VENDOR_BIT | BARRIER_BIT;
packets.stop_packet.header = VENDOR_BIT | BARRIER_BIT;
packets.start_packet.completion_signal = hsa_signal_t{.handle = 0};
packets.stop_packet.completion_signal = hsa_signal_t{.handle = 0};
this->empty = false;
@@ -147,8 +147,8 @@ QueueController::add_queue(hsa_queue_t* id, std::unique_ptr<Queue> queue)
{
for(const auto& itr : context::get_registered_contexts())
{
if(itr->thread_trace)
itr->thread_trace->resource_init(queue->get_agent(), get_core_table(), get_ext_table());
if(auto* trace = itr->dispatch_thread_trace.get())
trace->resource_init(queue->get_agent(), get_core_table(), get_ext_table());
}
CHECK(queue);
@@ -175,11 +175,11 @@ QueueController::destroy_queue(hsa_queue_t* id)
for(const auto& itr : context::get_registered_contexts())
{
if(!itr->thread_trace) continue;
if(!itr->dispatch_thread_trace) continue;
_queues.wlock([&](auto& map) {
if(map.find(id) != map.end())
itr->thread_trace->resource_deinit(map.at(id)->get_agent());
itr->dispatch_thread_trace->resource_deinit(map.at(id)->get_agent());
});
}
@@ -264,7 +264,7 @@ QueueController::init(CoreApiTable& core_table, AmdExtTable& ext_table)
auto enable_intercepter = false;
for(const auto& itr : context::get_registered_contexts())
{
constexpr auto expected_context_size = 208UL;
constexpr auto expected_context_size = 216UL;
static_assert(
sizeof(context::context) == expected_context_size,
"If you added a new field to context struct, make sure there is a check here if it "
@@ -277,7 +277,7 @@ QueueController::init(CoreApiTable& core_table, AmdExtTable& ext_table)
itr->buffered_tracer->domains(ROCPROFILER_BUFFER_TRACING_KERNEL_DISPATCH));
if(itr->counter_collection || itr->pc_sampler || has_kernel_tracing ||
itr->agent_counter_collection || itr->thread_trace)
itr->agent_counter_collection || itr->agent_thread_trace || itr->dispatch_thread_trace)
{
enable_intercepter = true;
break;
+2 -1
View File
@@ -631,6 +631,7 @@ finalize()
set_fini_status(-1);
hsa::async_copy_fini();
hsa::queue_controller_fini();
thread_trace::finalize();
page_migration::finalize();
#if ROCPROFILER_SDK_HSA_PC_SAMPLING > 0
// WARNING: this must precede `code_object::finalize()`
@@ -788,7 +789,7 @@ rocprofiler_set_api_table(const char* name,
rocprofiler::hsa::async_copy_init(hsa_api_table, lib_instance);
rocprofiler::code_object::initialize(hsa_api_table);
rocprofiler::thread_trace::code_object::initialize(hsa_api_table);
rocprofiler::thread_trace::initialize(hsa_api_table);
#if ROCPROFILER_SDK_HSA_PC_SAMPLING > 0
rocprofiler::pc_sampling::code_object::initialize(hsa_api_table);
#endif
@@ -51,12 +51,14 @@
} \
}
constexpr size_t ROCPROFILER_QUEUE_SIZE = 64;
namespace rocprofiler
{
namespace thread_trace
{
constexpr size_t QUEUE_SIZE = 128;
constexpr uint64_t MIN_BUFFER_SIZE = 1 << 18; // 2 pages per SE
constexpr uint64_t MAX_BUFFER_SIZE = std::numeric_limits<int32_t>::max(); // aqlprofile limit
struct cbdata_t
{
rocprofiler_att_shader_data_callback_t cb_fn;
@@ -65,10 +67,78 @@ struct cbdata_t
common::Synchronized<std::optional<int64_t>> client;
bool
ThreadTracerQueue::Submit(hsa_ext_amd_aql_pm4_packet_t* packet)
CoreApiTable&
get_core()
{
const uint64_t write_idx = add_write_index_relaxed_fn(queue, 1);
static CoreApiTable api{};
return api;
}
AmdExtTable&
get_ext()
{
static AmdExtTable api{};
return api;
}
bool
thread_trace_parameter_pack::are_params_valid() const
{
if(shader_cb_fn == nullptr)
{
ROCP_WARNING << "Callback cannot be null!";
return false;
}
if(shader_engine_mask == 0) return false;
if(buffer_size > MAX_BUFFER_SIZE || buffer_size < MIN_BUFFER_SIZE)
{
ROCP_WARNING << "Invalid buffer size: " << buffer_size;
return false;
}
if(target_cu > 0xF) return false;
if(simd_select > 0xF) return false; // Only 16 CUs and 4 SIMDs
return true;
}
class Signal
{
public:
Signal(hsa_ext_amd_aql_pm4_packet_t* packet)
{
get_ext().hsa_amd_signal_create_fn(0, 0, nullptr, 0, &signal);
packet->completion_signal = signal;
get_core().hsa_signal_store_screlease_fn(signal, 1);
};
~Signal()
{
WaitOn();
get_core().hsa_signal_destroy_fn(signal);
}
Signal(Signal& other) = delete;
Signal(const Signal& other) = delete;
Signal& operator=(Signal& other) = delete;
Signal& operator=(const Signal& other) = delete;
void WaitOn() const
{
auto* wait_fn = get_core().hsa_signal_wait_scacquire_fn;
while(wait_fn(signal, HSA_SIGNAL_CONDITION_EQ, 0, UINT64_MAX, HSA_WAIT_STATE_BLOCKED))
;
}
hsa_signal_t signal;
std::atomic<bool> released{false};
};
std::unique_ptr<Signal>
ThreadTracerQueue::Submit(hsa_ext_amd_aql_pm4_packet_t* packet, bool bWait)
{
std::unique_ptr<Signal> signal{};
const uint64_t write_idx = add_write_index_relaxed_fn(queue, 1);
size_t index = (write_idx % queue->size) * sizeof(hsa_ext_amd_aql_pm4_packet_t);
auto* queue_slot = reinterpret_cast<uint32_t*>(size_t(queue->base_address) + index); // NOLINT
@@ -76,23 +146,15 @@ ThreadTracerQueue::Submit(hsa_ext_amd_aql_pm4_packet_t* packet)
const auto* slot_data = reinterpret_cast<const uint32_t*>(packet);
memcpy(&queue_slot[1], &slot_data[1], sizeof(hsa_ext_amd_aql_pm4_packet_t) - sizeof(uint32_t));
if(bWait)
signal =
std::make_unique<Signal>(reinterpret_cast<hsa_ext_amd_aql_pm4_packet_t*>(queue_slot));
auto* header = reinterpret_cast<std::atomic<uint32_t>*>(queue_slot);
header->store(slot_data[0], std::memory_order_release);
signal_store_screlease_fn(queue->doorbell_signal, write_idx);
int loops = 0;
while(load_read_index_relaxed_fn(queue) <= write_idx)
{
loops++;
usleep(1);
if(loops > 10000) // Add loop limit to prevent hang. TODO: Remove once stability proven
{
ROCP_ERROR << "Codeobj packet submission failed!";
return false;
}
}
return true;
return signal;
}
ThreadTracerQueue::ThreadTracerQueue(thread_trace_parameter_pack _params,
@@ -106,7 +168,7 @@ ThreadTracerQueue::ThreadTracerQueue(thread_trace_parameter_pack _params,
control_packet = factory->construct_control_packet();
auto status = coreapi.hsa_queue_create_fn(cache.get_hsa_agent(),
ROCPROFILER_QUEUE_SIZE,
QUEUE_SIZE,
HSA_QUEUE_TYPE_SINGLE,
nullptr,
nullptr,
@@ -147,8 +209,10 @@ ThreadTracerQueue::~ThreadTracerQueue()
control_packet->clear();
control_packet->populate_after();
std::vector<std::unique_ptr<Signal>> wait_idx{};
for(auto& after_packet : control_packet->after_krn_pkt)
Submit(&after_packet);
wait_idx.emplace_back(Submit(&after_packet, true));
}
/**
@@ -200,11 +264,8 @@ ThreadTracerQueue::load_codeobj(code_object_id_t id, uint64_t addr, uint64_t siz
if(!queue || active_traces.load() < 1) return;
auto packet = factory->construct_load_marker_packet(id, addr, size);
bool bSuccess = Submit(&packet->packet);
if(!bSuccess) // If something went wrong, don't delete packet to avoid CP memory access fault
packet.release();
auto packet = factory->construct_load_marker_packet(id, addr, size);
Submit(&packet->packet, true)->WaitOn();
}
void
@@ -215,11 +276,8 @@ ThreadTracerQueue::unload_codeobj(code_object_id_t id)
if(!control_packet->remove_codeobj(id)) return;
if(!queue || active_traces.load() < 1) return;
auto packet = factory->construct_unload_marker_packet(id);
bool bSuccess = Submit(&packet->packet);
if(!bSuccess) // If something went wrong, don't delete packet to avoid CP memory access fault
packet.release();
auto packet = factory->construct_unload_marker_packet(id);
Submit(&packet->packet, true)->WaitOn();
}
void
@@ -301,15 +359,13 @@ DispatchThreadTracer::pre_kernel_call(const hsa::Queue& queue,
auto it = agents.find(queue.get_agent().get_hsa_agent());
assert(it != agents.end() && it->second != nullptr);
auto packet = it->second->get_control(bool(control_flags & ROCPROFILER_ATT_CONTROL_START));
auto packet = it->second->get_control(true);
post_move_data.fetch_add(1);
maybe_add_serialization(packet);
if((control_flags & ROCPROFILER_ATT_CONTROL_START) != 0) packet->populate_before();
if((control_flags & ROCPROFILER_ATT_CONTROL_STOP) != 0) packet->populate_after();
packet->populate_before();
packet->populate_after();
return packet;
}
@@ -383,7 +439,7 @@ DispatchThreadTracer::start_context()
}
void
DispatchThreadTracer::stop_context()
DispatchThreadTracer::stop_context() // NOLINT
{
client.wlock([&](auto& client_id) {
if(!client_id) return;
@@ -398,34 +454,33 @@ DispatchThreadTracer::stop_context()
}
void
AgentThreadTracer::resource_init(const hsa::AgentCache& cache,
const CoreApiTable& coreapi,
const AmdExtTable& ext)
AgentThreadTracer::resource_init(const CoreApiTable& coreapi, const AmdExtTable& ext)
{
auto id = cache.get_rocp_agent()->id;
auto rocp_agents = rocprofiler::agent::get_agents();
std::unique_lock<std::mutex> lk(agent_mut);
if(params.find(id) == params.end()) return;
if(tracers.find(id) != tracers.end())
for(const auto* rocp_agent : rocp_agents)
{
tracers.at(id)->active_queues.fetch_add(1);
return;
auto id = rocp_agent->id;
const auto* cache = rocprofiler::agent::get_agent_cache(rocp_agent);
if(tracers.find(id) != tracers.end())
ROCP_WARNING << "Agent configured twice: " << id.handle;
else if(params.find(id) == params.end())
ROCP_INFO << "Skipping agent " << id.handle;
else if(cache == nullptr)
ROCP_WARNING << "Invalid agent id: " << id.handle;
else
tracers[id] = std::make_unique<ThreadTracerQueue>(params.at(id), *cache, coreapi, ext);
}
tracers.emplace(id, std::make_unique<ThreadTracerQueue>(params.at(id), cache, coreapi, ext));
}
void
AgentThreadTracer::resource_deinit(const hsa::AgentCache& cache)
AgentThreadTracer::resource_deinit()
{
auto id = cache.get_rocp_agent()->id;
std::unique_lock<std::mutex> lk(agent_mut);
if(params.find(id) == params.end()) return;
if(tracers.find(id) == tracers.end()) return;
auto& tracer = *tracers.at(id);
if(tracer.active_queues.fetch_sub(1) == 1) tracers.erase(id);
tracers.clear();
}
void
@@ -439,13 +494,15 @@ AgentThreadTracer::start_context()
return;
}
std::vector<std::unique_ptr<Signal>> wait_list{};
for(auto& [_, tracer] : tracers)
{
auto packet = tracer->get_control(true);
packet->populate_before();
for(auto& start : packet->before_krn_pkt)
tracer->Submit(&start);
auto sig = tracer->SubmitAndSignalLast(packet->before_krn_pkt);
if(sig) wait_list.emplace_back(std::move(sig));
}
}
@@ -460,16 +517,50 @@ AgentThreadTracer::stop_context()
return;
}
for(auto& [_, tracer] : tracers)
std::vector<std::tuple<ThreadTracerQueue*, aqlprofile_handle_t, std::unique_ptr<Signal>>>
wait_list{};
for(auto& [id, tracer] : tracers)
{
auto packet = tracer->get_control(false);
packet->populate_after();
for(auto& stop : packet->after_krn_pkt)
tracer->Submit(&stop);
std::optional<uint64_t> write_index{};
auto signal = tracer->SubmitAndSignalLast(packet->after_krn_pkt);
if(signal)
wait_list.push_back({tracer.get(), packet->GetHandle(), std::move(signal)}); // NOLINT
}
for(auto& [tracer, handle, signal] : wait_list)
{
signal->WaitOn();
rocprofiler_user_data_t userdata{.ptr = tracer->params.callback_userdata};
tracer->iterate_data(packet->GetHandle(), userdata);
tracer->iterate_data(handle, userdata);
}
}
void
initialize(HsaApiTable* table)
{
assert(table->core_ && table->amd_ext_);
get_core() = *table->core_;
get_ext() = *table->amd_ext_;
code_object::initialize(table);
for(auto& ctx : context::get_registered_contexts())
{
if(ctx->agent_thread_trace)
ctx->agent_thread_trace->resource_init(*table->core_, *table->amd_ext_);
}
}
void
finalize()
{
for(auto& ctx : context::get_registered_contexts())
{
if(ctx->agent_thread_trace) ctx->agent_thread_trace->resource_deinit();
}
}
@@ -71,6 +71,8 @@ struct thread_trace_parameter_pack
static constexpr size_t DEFAULT_SE_MASK = 0x21;
static constexpr size_t DEFAULT_BUFFER_SIZE = 0x8000000;
static constexpr size_t PERFCOUNTER_SIMD_MASK_SHIFT = 28;
bool are_params_valid() const;
};
class ThreadTracerQueue
@@ -99,7 +101,19 @@ public:
std::unique_ptr<hsa::TraceControlAQLPacket> control_packet;
std::unique_ptr<aql::ThreadTraceAQLPacketFactory> factory;
bool Submit(hsa_ext_amd_aql_pm4_packet_t* packet);
[[nodiscard]] std::unique_ptr<class Signal> Submit(hsa_ext_amd_aql_pm4_packet_t* packet,
bool bWait);
template <typename VecType>
[[nodiscard]] std::unique_ptr<class Signal> SubmitAndSignalLast(VecType vec)
{
for(size_t i = 0; i < vec.size(); i++)
{
auto sig = Submit(&vec.at(i), i == vec.size() - 1);
if(sig) return sig;
}
return nullptr;
}
private:
std::unique_ptr<code_object::CodeobjCallbackRegistry> codeobj_reg{nullptr};
@@ -112,19 +126,7 @@ private:
decltype(hsa_queue_destroy)* queue_destroy_fn{nullptr};
};
class ThreadTracerInterface
{
public:
ThreadTracerInterface() = default;
virtual ~ThreadTracerInterface() = default;
virtual void start_context() = 0;
virtual void stop_context() = 0;
virtual void resource_init(const hsa::AgentCache&, const CoreApiTable&, const AmdExtTable&) = 0;
virtual void resource_deinit(const hsa::AgentCache&) = 0;
};
class DispatchThreadTracer : public ThreadTracerInterface
class DispatchThreadTracer
{
using code_object_id_t = uint64_t;
using AQLPacketPtr = std::unique_ptr<hsa::AQLPacket>;
@@ -134,12 +136,12 @@ public:
DispatchThreadTracer(thread_trace_parameter_pack _params)
: params(std::move(_params))
{}
~DispatchThreadTracer() override = default;
~DispatchThreadTracer() = default;
void start_context() override;
void stop_context() override;
void resource_init(const hsa::AgentCache&, const CoreApiTable&, const AmdExtTable&) override;
void resource_deinit(const hsa::AgentCache&) override;
void start_context();
void stop_context();
void resource_init(const hsa::AgentCache&, const CoreApiTable&, const AmdExtTable&);
void resource_deinit(const hsa::AgentCache&);
std::unique_ptr<hsa::AQLPacket> pre_kernel_call(const hsa::Queue& queue,
uint64_t kernel_id,
@@ -157,16 +159,16 @@ public:
thread_trace_parameter_pack params;
};
class AgentThreadTracer : public ThreadTracerInterface
class AgentThreadTracer
{
public:
AgentThreadTracer() = default;
~AgentThreadTracer() override = default;
AgentThreadTracer() = default;
~AgentThreadTracer() = default;
void start_context() override;
void stop_context() override;
void resource_init(const hsa::AgentCache&, const CoreApiTable&, const AmdExtTable&) override;
void resource_deinit(const hsa::AgentCache&) override;
void start_context();
void stop_context();
void resource_init(const CoreApiTable&, const AmdExtTable&);
void resource_deinit();
void add_agent(rocprofiler_agent_id_t id, thread_trace_parameter_pack _params)
{
@@ -180,11 +182,17 @@ public:
}
std::map<rocprofiler_agent_id_t, std::unique_ptr<ThreadTracerQueue>> tracers{};
std::map<rocprofiler_agent_id_t, thread_trace_parameter_pack> params;
std::map<rocprofiler_agent_id_t, thread_trace_parameter_pack> params{};
std::mutex agent_mut;
};
void
initialize(HsaApiTable* table);
void
finalize();
}; // namespace thread_trace
} // namespace rocprofiler
@@ -30,6 +30,9 @@
#include "lib/rocprofiler-sdk/registration.hpp"
#include "rocprofiler-sdk/amd_detail/thread_trace.h"
using DispatchThreadTracer = rocprofiler::thread_trace::DispatchThreadTracer;
using AgentThreadTracer = rocprofiler::thread_trace::AgentThreadTracer;
extern "C" {
rocprofiler_status_t ROCPROFILER_API
rocprofiler_configure_dispatch_thread_trace_service(
@@ -45,7 +48,8 @@ rocprofiler_configure_dispatch_thread_trace_service(
auto* ctx = rocprofiler::context::get_mutable_registered_context(context_id);
if(!ctx) return ROCPROFILER_STATUS_ERROR_CONTEXT_NOT_STARTED;
if(ctx->thread_trace) return ROCPROFILER_STATUS_ERROR_SERVICE_ALREADY_CONFIGURED;
if(ctx->dispatch_thread_trace) return ROCPROFILER_STATUS_ERROR_SERVICE_ALREADY_CONFIGURED;
if(ctx->agent_thread_trace) return ROCPROFILER_STATUS_ERROR_CONTEXT_INVALID;
auto pack = rocprofiler::thread_trace::thread_trace_parameter_pack{};
@@ -54,6 +58,8 @@ rocprofiler_configure_dispatch_thread_trace_service(
pack.shader_cb_fn = shader_callback;
pack.callback_userdata = callback_userdata;
if(pack.dispatch_cb_fn == nullptr) return ROCPROFILER_STATUS_ERROR_INVALID_ARGUMENT;
auto id_map = rocprofiler::counters::getPerfCountersIdMap();
for(size_t p = 0; p < num_parameters; p++)
{
@@ -83,7 +89,9 @@ rocprofiler_configure_dispatch_thread_trace_service(
}
}
ctx->thread_trace = std::make_unique<rocprofiler::thread_trace::DispatchThreadTracer>(pack);
if(!pack.are_params_valid()) return ROCPROFILER_STATUS_ERROR_INVALID_ARGUMENT;
ctx->dispatch_thread_trace = std::make_unique<DispatchThreadTracer>(pack);
return ROCPROFILER_STATUS_SUCCESS;
}
@@ -96,14 +104,14 @@ rocprofiler_configure_agent_thread_trace_service(
rocprofiler_att_shader_data_callback_t shader_callback,
void* callback_userdata)
{
using AgentThreadTracer = rocprofiler::thread_trace::AgentThreadTracer;
if(rocprofiler::registration::get_init_status() > -1)
return ROCPROFILER_STATUS_ERROR_CONFIGURATION_LOCKED;
auto* ctx = rocprofiler::context::get_mutable_registered_context(context_id);
if(!ctx) return ROCPROFILER_STATUS_ERROR_CONTEXT_NOT_STARTED;
if(ctx->dispatch_thread_trace) return ROCPROFILER_STATUS_ERROR_CONTEXT_INVALID;
if(!ctx->thread_trace) ctx->thread_trace = std::make_unique<AgentThreadTracer>();
if(!ctx->agent_thread_trace) ctx->agent_thread_trace = std::make_unique<AgentThreadTracer>();
auto pack = rocprofiler::thread_trace::thread_trace_parameter_pack{};
@@ -140,11 +148,10 @@ rocprofiler_configure_agent_thread_trace_service(
}
}
auto* agent_tracer = dynamic_cast<AgentThreadTracer*>(ctx->thread_trace.get());
if(agent_tracer == nullptr || agent_tracer->has_agent(agent))
return ROCPROFILER_STATUS_ERROR_SERVICE_ALREADY_CONFIGURED;
if(!pack.are_params_valid()) return ROCPROFILER_STATUS_ERROR_INVALID_ARGUMENT;
agent_tracer->add_agent(agent, pack);
assert(ctx->agent_thread_trace);
ctx->agent_thread_trace->add_agent(agent, pack);
return ROCPROFILER_STATUS_SUCCESS;
}
}
@@ -209,7 +209,7 @@ TEST(thread_trace, perfcounters_configure_test)
nullptr);
auto* context = rocprofiler::context::get_mutable_registered_context(ctx);
auto* tracer = dynamic_cast<thread_trace::DispatchThreadTracer*>(context->thread_trace.get());
auto* tracer = context->dispatch_thread_trace.get();
ASSERT_NE(tracer, nullptr);
ASSERT_EQ(tracer->params.perfcounter_ctrl, 1);
@@ -250,3 +250,61 @@ TEST(thread_trace, perfcounters_aql_options_test)
context::pop_client(1);
hsa_shut_down();
}
rocprofiler_status_t
query_available_agents(rocprofiler_agent_version_t /* version */,
const void** agents,
size_t num_agents,
void* ctx_ptr)
{
for(size_t idx = 0; idx < num_agents; idx++)
{
const auto* agent = static_cast<const rocprofiler_agent_v0_t*>(agents[idx]);
if(agent->type != ROCPROFILER_AGENT_TYPE_GPU) continue;
std::vector<rocprofiler_att_parameter_t> params;
params.push_back({ROCPROFILER_ATT_PARAMETER_TARGET_CU, {1}});
params.push_back({ROCPROFILER_ATT_PARAMETER_SHADER_ENGINE_MASK, {0xF}});
params.push_back({ROCPROFILER_ATT_PARAMETER_BUFFER_SIZE, {0x1000000}});
params.push_back({ROCPROFILER_ATT_PARAMETER_SIMD_SELECT, {0xF}});
params.push_back({ROCPROFILER_ATT_PARAMETER_PERFCOUNTERS_CTRL, {1}});
{
auto metrics = rocprofiler::counters::getMetricsForAgent("gfx90a");
rocprofiler_att_parameter_t att_param;
att_param.type = ROCPROFILER_ATT_PARAMETER_PERFCOUNTER;
att_param.simd_mask = 0xF;
for(auto& metric : metrics)
if(metric.name() == "SQ_WAVES") rocprofiler_counter_id_t{.handle = metric.id()};
params.push_back(att_param);
}
rocprofiler_configure_agent_thread_trace_service(
*reinterpret_cast<rocprofiler_context_id_t*>(ctx_ptr),
params.data(),
params.size(),
agent->id,
[](int64_t, void*, size_t, rocprofiler_user_data_t) {},
nullptr);
}
return ROCPROFILER_STATUS_SUCCESS;
}
TEST(thread_trace, agent_configure_test)
{
test_init();
registration::init_logging();
registration::set_init_status(-1);
context::push_client(1);
rocprofiler_context_id_t ctx;
ROCPROFILER_CALL(rocprofiler_create_context(&ctx), "context creation failed");
ROCPROFILER_CALL(rocprofiler_query_available_agents(ROCPROFILER_AGENT_INFO_VERSION_0,
&query_available_agents,
sizeof(rocprofiler_agent_t),
&ctx),
"Failed to find GPU agents");
}
+3 -22
View File
@@ -54,37 +54,18 @@ rocprofiler_att_control_flags_t
dispatch_callback(rocprofiler_queue_id_t /* queue_id */,
const rocprofiler_agent_t* /* agent */,
rocprofiler_correlation_id_t /* correlation_id */,
rocprofiler_kernel_id_t kernel_id,
rocprofiler_kernel_id_t /* kernel_id */,
rocprofiler_dispatch_id_t /* dispatch_id */,
rocprofiler_user_data_t* dispatch_userdata,
void* userdata)
{
C_API_BEGIN
assert(userdata && "Dispatch callback passed null!");
ToolData& tool = *reinterpret_cast<ToolData*>(userdata);
dispatch_userdata->ptr = userdata;
static std::atomic<int> call_id{0};
static std::string_view desired_func_name = "branching_kernel";
try
{
auto& kernel_name = tool.kernel_id_to_kernel_name.at(kernel_id);
if(kernel_name.find(desired_func_name) == std::string::npos)
return ROCPROFILER_ATT_CONTROL_NONE;
int id = call_id.fetch_add(1);
if(id == 0)
return ROCPROFILER_ATT_CONTROL_START;
else if(id == 1)
return ROCPROFILER_ATT_CONTROL_STOP;
} catch(...)
{
std::cerr << "Could not find kernel id: " << kernel_id << std::endl;
}
C_API_END
return ROCPROFILER_ATT_CONTROL_NONE;
return ROCPROFILER_ATT_CONTROL_START_AND_STOP;
}
int