diff --git a/docs/README_clocksync.md b/docs/README_clocksync.md new file mode 100644 index 0000000..e6e6e15 --- /dev/null +++ b/docs/README_clocksync.md @@ -0,0 +1,60 @@ +# ChronoSync: Multi-Node Clock Synchronization + +ChronoSync enables accurate cross-node timestamp correlation for distributed GPU workloads. It uses the Firefly clock synchronization algorithm to measure and correct clock offset and drift between nodes, so that traces from multiple machines can be merged and compared with consistent timing. + +## How it works + +When enabled, one process per node runs the Firefly protocol — exchanging UDP timing probes with peer nodes to compute clock offset and drift via linear regression. The correction is applied to all timestamps (CPU-side API calls and GPU activity from roctracer/rocprofiler-sdk/cupti) via a shared memory region that all profiled processes on the node can read. + +## Configuration + +Two environment variables control ChronoSync: + +### `RPDT_CLOCKSYNC_IP` + +Path to a configuration file listing all participating nodes. Each line has the format: + +``` +,rank= +``` + +Example config file for a two-node setup: + +``` +192.0.2.1,rank=0 +192.0.2.2,rank=1 +``` + +### `RPDT_CLOCKSYNC_RANK` + +The rank of the current node. Must match one of the ranks in the config file. + +## Usage + +1. Create a config file listing all nodes and their ranks. + +2. On each node, set the environment variables and run the workload with rpd_tracer: + +```bash +# Node 0 +export RPDT_CLOCKSYNC_IP=/path/to/sync_config.txt +export RPDT_CLOCKSYNC_RANK=0 +LD_PRELOAD=librpd_tracer.so ./my_workload + +# Node 1 +export RPDT_CLOCKSYNC_IP=/path/to/sync_config.txt +export RPDT_CLOCKSYNC_RANK=1 +LD_PRELOAD=librpd_tracer.so ./my_workload +``` + +3. All processes on each node that write to the same RPD file automatically share the clock correction via POSIX shared memory. No additional configuration is needed for multi-process workloads. + +## Network requirements + +- Nodes must be able to reach each other via UDP and TCP on port `12345 + rank_a + rank_b` for each pair of nodes. +- The Firefly protocol uses a TCP handshake for initial synchronization, then UDP probes for ongoing measurement. + +## Limitations + +- Clock correction degrades over time after the sync process exits. For long-running sequential workloads, keep the sync process active for the duration of the trace. +- The drift rate is clamped at 500 ppm. If the measured drift exceeds this threshold (indicating noisy measurements), drift correction is disabled and only the offset is applied. diff --git a/rpd_tracer/ChronoSyncDataSource.cpp b/rpd_tracer/ChronoSyncDataSource.cpp new file mode 100644 index 0000000..41634e3 --- /dev/null +++ b/rpd_tracer/ChronoSyncDataSource.cpp @@ -0,0 +1,357 @@ +/************************************************************************** + * Copyright (c) 2022 Advanced Micro Devices, Inc. + **************************************************************************/ +#include "ChronoSyncDataSource.h" + +#include +#include +#include +#include +#include +#include + +#include +#include +#include +#include +#include +#include +#include + +#include + +#include "DbResource.h" +#include "Logger.h" +#include "Utility.h" +#include "Firefly.h" + +using rpdtracer::DataSource; +using rpdtracer::ChronoSyncDataSource; + +// Create a factory for the Logger to locate and use +extern "C" { +DataSource* ChronoSyncDataSourceFactory() { + return new ChronoSyncDataSource(); +} +} + +namespace rpdtracer { + +// ----------------------------------------------------------------------------- +// ChronoSyncDataSourcePrivate +// ----------------------------------------------------------------------------- +class ChronoSyncDataSourcePrivate { +public: + explicit ChronoSyncDataSourcePrivate(ChronoSyncDataSource* owner) + : m_owner(owner) {} + + void work() { + if (m_owner == nullptr) { +// std::fprintf(stderr, "ChronoSync: [ChronoSyncDataSourcePrivate::work] CRITICAL - Owner not set\n"); + return; + } + + std::unique_lock lock(m_owner->m_mutex); + while (!m_done.load(std::memory_order_relaxed)) { + if (!m_owner->m_workExecuted) { +// std::fprintf(stderr, "ChronoSync: [ChronoSync Worker] INFO - Executing work once\n"); + m_workerRunning = true; + lock.unlock(); + m_owner->work(); + lock.lock(); + m_owner->m_workExecuted = true; + m_workerRunning = false; +// std::fprintf(stderr, "ChronoSync: [ChronoSync Worker] INFO - Work completed\n"); + } + + m_workerRunning = false; + if (!m_done.load(std::memory_order_relaxed)) { +// std::fprintf(stderr, "ChronoSync: [ChronoSync Worker] INFO - Waiting for signal\n"); + m_owner->m_wait.wait(lock); + } + m_workerRunning = true; + } + +// std::fprintf(stderr, "ChronoSync: [ChronoSync Worker] INFO - Exiting\n"); + } + + ChronoSyncDataSource* m_owner{nullptr}; + std::thread* m_worker{nullptr}; + std::atomic m_done{false}; + bool m_workerRunning{false}; + std::string m_hostIp; + int m_rank{-1}; + std::vector> m_neighbors; +}; + +// ----------------------------------------------------------------------------- +// ChronoSyncDataSource Implementation +// ----------------------------------------------------------------------------- +ChronoSyncDataSource::ChronoSyncDataSource() + : m_private(nullptr), + m_resource(nullptr), + m_workExecuted(false), + m_messageCount(0) {} + +ChronoSyncDataSource::~ChronoSyncDataSource() { + end(); + delete m_private; + m_private = nullptr; + delete m_resource; + m_resource = nullptr; +} + +static std::string rpd_filename() +{ + const char *f = getenv("RPDT_FILENAME"); + return (f != nullptr) ? f : "./trace.rpd"; +} + +static int metadata_callback(void *data, int argc, char **argv, char **colName) +{ + std::string &value = *static_cast(data); + if (argc > 0 && argv[0]) + value = argv[0]; + return 0; +} + +void ChronoSyncDataSource::storeMetadata(const std::string& tag, const std::string& value) +{ + sqlite3 *db = nullptr; + if (sqlite3_open(rpd_filename().c_str(), &db) != SQLITE_OK) + return; + sqlite3_busy_handler(db, &sqlite_busy_handler, NULL); + char *err = nullptr; + std::string sql = "INSERT OR REPLACE INTO rocpd_metadata(tag, value) VALUES ('" + tag + "', '" + value + "')"; + sqlite3_exec(db, sql.c_str(), nullptr, nullptr, &err); + sqlite3_close(db); +} + +std::string ChronoSyncDataSource::queryMetadata(const std::string& tag) +{ + std::string value; + sqlite3 *db = nullptr; + if (sqlite3_open(rpd_filename().c_str(), &db) != SQLITE_OK) + return value; + sqlite3_busy_handler(db, &sqlite_busy_handler, NULL); + char *err = nullptr; + std::string sql = "SELECT value FROM rocpd_metadata WHERE tag = '" + tag + "'"; + sqlite3_exec(db, sql.c_str(), metadata_callback, &value, &err); + sqlite3_close(db); + return value; +} + +void ChronoSyncDataSource::init() { +// std::fprintf(stderr, "ChronoSync: [ChronoSyncDataSource::init] INFO - Called (PID: %d)\n", getpid()); + + if (m_private != nullptr) { +// std::fprintf(stderr, "ChronoSync: [ChronoSyncDataSource::init] WARNING - Already initialized\n"); + return; + } + + m_resource = new DbResource(rpd_filename(), std::string("chronosync_active")); + if (!m_resource->tryLock()) { +// std::fprintf(stderr, "ChronoSync: [ChronoSyncDataSource::init] WARNING - Another instance active\n"); + // Try to attach to existing shared memory from the singleton + std::string existingShm = queryMetadata("clocksync_shm"); + if (!existingShm.empty()) + firefly::attach_clocksync_shm(existingShm); + return; + } + +// std::fprintf(stderr, "ChronoSync: [ChronoSyncDataSource::init] INFO - Lock acquired (PID: %d)\n", getpid()); + + firefly::create_clocksync_shm(m_shmName); + storeMetadata("clocksync_shm", m_shmName); + + m_private = new ChronoSyncDataSourcePrivate(this); + m_private->m_workerRunning = true; + m_private->m_worker = new std::thread(&ChronoSyncDataSourcePrivate::work, m_private); + +// std::fprintf(stderr, "ChronoSync: [ChronoSyncDataSource::init] INFO - Worker thread created\n"); +} + +void ChronoSyncDataSource::startTracing() { +// std::fprintf(stderr, "ChronoSync: [ChronoSyncDataSource::startTracing] INFO - Called (PID: %d)\n", getpid()); + + if (m_private == nullptr) { +// std::fprintf(stderr, "ChronoSync: [ChronoSyncDataSource::startTracing] WARNING - Not singleton instance\n"); + return; + } + + if (m_workExecuted) { +// std::fprintf(stderr, "ChronoSync: [ChronoSyncDataSource::startTracing] WARNING - Work already executed\n"); + return; + } + + std::lock_guard lock(m_mutex); + m_wait.notify_all(); +} + +void ChronoSyncDataSource::stopTracing() { +// std::fprintf(stderr, "ChronoSync: [ChronoSyncDataSource::stopTracing] INFO - Called (PID: %d)\n", getpid()); + if (m_private == nullptr) { +// std::fprintf(stderr, "ChronoSync: [ChronoSyncDataSource::stopTracing] WARNING - Not singleton instance\n"); + } +} + +void ChronoSyncDataSource::flush() { + if (m_private == nullptr) { + return; + } + +// std::fprintf(stderr, "ChronoSync: [ChronoSyncDataSource::flush] INFO - Called\n"); + std::unique_lock lock(m_mutex); + while (m_private->m_workerRunning) { + m_wait.wait(lock); + } +} + +void ChronoSyncDataSource::end() { +// std::fprintf(stderr, "ChronoSync: [ChronoSyncDataSource::end] INFO - Called (PID: %d)\n", getpid()); + + if (m_private == nullptr) + return; + + if (m_private->m_worker != nullptr) { + m_private->m_done.store(true, std::memory_order_relaxed); + m_wait.notify_one(); + m_private->m_worker->join(); + delete m_private->m_worker; + m_private->m_worker = nullptr; + } + + if (m_resource != nullptr) + m_resource->unlock(); + + if (!m_shmName.empty()) + firefly::cleanup_clocksync_shm(m_shmName); + +// std::fprintf(stderr, "ChronoSync: [ChronoSyncDataSource::end] INFO - Cleanup complete\n"); +} + +// ----------------------------------------------------------------------------- +// Work Routine +// ----------------------------------------------------------------------------- +void ChronoSyncDataSource::work() { +// std::fprintf(stderr, "ChronoSync: [ChronoSyncDataSource::work] INFO - Starting\n"); + ++m_messageCount; + + auto now = std::chrono::system_clock::now(); + const std::time_t nowSeconds = std::chrono::system_clock::to_time_t(now); + auto nowMilliseconds = std::chrono::duration_cast(now.time_since_epoch()) % 1000; + char timeBuffer[26]; + ctime_r(&nowSeconds, timeBuffer); + timeBuffer[24] = '\0'; +// std::fprintf(stderr, +// "[%s.%03lld] ChronoSync Work #%d: Performing synchronization...\n", +// timeBuffer, +// static_cast(nowMilliseconds.count()), +// m_messageCount); + + const char* configPath = std::getenv("RPDT_CLOCKSYNC_IP"); + if (configPath == nullptr) { +// std::fprintf(stderr, "ChronoSync: [ChronoSyncDataSource::work] WARNING - RPDT_CLOCKSYNC_IP not set\n"); + return; + } + + FILE* fileHandle = std::fopen(configPath, "r"); + if (fileHandle == nullptr) { +// std::fprintf(stderr, +// "ChronoSync: [ChronoSyncDataSource::work] CRITICAL - Failed to open %s (errno: %s)\n", +// configPath, +// std::strerror(errno)); + return; + } + + std::string hostIp; + int hostRank = -1; + std::vector> neighbors; + char lineBuffer[256]; + + const char* myRankEnv = std::getenv("RPDT_CLOCKSYNC_RANK"); + int myRank = myRankEnv ? std::stoi(myRankEnv) : -1; + + while (std::fgets(lineBuffer, sizeof(lineBuffer), fileHandle)) { + std::string line(lineBuffer); + + // Trim whitespace + line.erase(0, line.find_first_not_of(" \t\n\r")); + line.erase(line.find_last_not_of(" \t\n\r") + 1); + + if (line.empty()) continue; + + // Parse: "10.7.76.147,rank=0" + size_t commaPos = line.find(','); + if (commaPos == std::string::npos) continue; + + std::string ip = line.substr(0, commaPos); + std::string rankStr = line.substr(commaPos + 1); + + // Extract rank value from "rank=0" + if (rankStr.substr(0, 5) != "rank=") continue; + + int fileRank = std::stoi(rankStr.substr(5)); + + if (myRank == fileRank) { + hostRank = fileRank; + hostIp = ip; + } else { + neighbors.emplace_back(ip, fileRank); + } + } + + std::fclose(fileHandle); + // print parsed values +// std::fprintf(stderr, "ChronoSync: [ChronoSyncDataSource::work] INFO - Host IP: %s, Host Rank: %d\n", hostIp.c_str(), hostRank); + for (const auto& neighbor : neighbors) { +// std::fprintf(stderr, "ChronoSync: [ChronoSyncDataSource::work] INFO - Neighbor IP: %s, Neighbor Rank: %d\n", neighbor.first.c_str(), neighbor.second); + } + + if (m_private == nullptr) { +// std::fprintf(stderr, "ChronoSync: [ChronoSyncDataSource::work] CRITICAL - Private data missing\n"); + return; + } + + m_private->m_hostIp = hostIp; + m_private->m_rank = hostRank; + m_private->m_neighbors = neighbors; + + firefly::MeasurementBuffer buffer(firefly::MAX_MEASUREMENT_COUNT); + bool done = false; + + std::vector workers; + for (const auto& neighbor : neighbors) { + const char* role = (hostRank < neighbor.second) ? "A" : "B"; + const char* nodeIpA = (hostRank < neighbor.second) ? hostIp.c_str() : neighbor.first.c_str(); + const char* nodeIpB = (hostRank < neighbor.second) ? neighbor.first.c_str() : hostIp.c_str(); + int udpPort = firefly::UDP_PORT_DEFAULT + hostRank + neighbor.second; + + std::string ipACopy(nodeIpA); + std::string ipBCopy(nodeIpB); + std::string roleCopy(role); + + workers.emplace_back([ipACopy, ipBCopy, roleCopy, udpPort, &buffer, &done]() { + firefly::run_node(roleCopy.c_str(), ipACopy.c_str(), ipBCopy.c_str(), + udpPort, buffer, done); + }); + } + + int lastCount = 0; + while (!m_private->m_done.load(std::memory_order_relaxed)) { + if (!neighbors.empty() && buffer.count != lastCount) { + lastCount = buffer.count; + const char* role = (hostRank < neighbors.front().second) ? "A" : "B"; + firefly::firefly_run(role, buffer); + } + std::this_thread::sleep_for(std::chrono::milliseconds(firefly::FIRE_FLY_SLEEP_MSEC)); + } + + done = true; + for (auto& worker : workers) + worker.join(); + +// std::fprintf(stderr, "ChronoSync: [ChronoSyncDataSource::work] INFO - Synchronization complete\n"); +} + +} // namespace rpdtracer diff --git a/rpd_tracer/ChronoSyncDataSource.h b/rpd_tracer/ChronoSyncDataSource.h new file mode 100644 index 0000000..13c4c02 --- /dev/null +++ b/rpd_tracer/ChronoSyncDataSource.h @@ -0,0 +1,51 @@ +/************************************************************************** + * Copyright (c) 2022 Advanced Micro Devices, Inc. + **************************************************************************/ +#ifndef CHRONOSYNCDATASOURCE_H +#define CHRONOSYNCDATASOURCE_H + +#include "DataSource.h" + +#include +#include +#include + +namespace rpdtracer { + +// Forward declarations +class ChronoSyncDataSourcePrivate; +class DbResource; + +class ChronoSyncDataSource : public DataSource +{ +public: + ChronoSyncDataSource(); + ~ChronoSyncDataSource(); + + void init() override; + void startTracing() override; + void stopTracing() override; + void flush() override; + void end() override; + + void work(); + + std::mutex m_mutex; + std::condition_variable m_wait; + bool m_workExecuted; + +private: + void storeMetadata(const std::string& tag, const std::string& value); + std::string queryMetadata(const std::string& tag); + + ChronoSyncDataSourcePrivate* m_private; + DbResource* m_resource; + std::string m_shmName; + int m_messageCount; + + friend class ChronoSyncDataSourcePrivate; +}; + +} // namespace rpdtracer + +#endif // CHRONOSYNCDATASOURCE_H diff --git a/rpd_tracer/CuptiDataSource.cpp b/rpd_tracer/CuptiDataSource.cpp index dfcb666..9d8de78 100644 --- a/rpd_tracer/CuptiDataSource.cpp +++ b/rpd_tracer/CuptiDataSource.cpp @@ -678,8 +678,8 @@ void CUPTIAPI CuptiDataSource::bufferCompleted(CUcontext ctx, uint32_t streamId, row.gpuId = record->deviceId; row.queueId = record->contextId; // FIXME: this or stream row.sequenceId = record->streamId; - row.start = record->start + toffset; - row.end = record->end + toffset; + row.start = adjust_external_ts(record->start + toffset); + row.end = adjust_external_ts(record->end + toffset); row.description_id = EMPTY_STRING_ID; row.opType_id = logger.stringTable().getOrCreate("Memcpy"); row.api_id = record->correlationId; @@ -692,8 +692,8 @@ void CUPTIAPI CuptiDataSource::bufferCompleted(CUcontext ctx, uint32_t streamId, row.gpuId = record->deviceId; row.queueId = record->contextId; // FIXME: this or stream row.sequenceId = record->streamId; - row.start = record->start + toffset; - row.end = record->end + toffset; + row.start = adjust_external_ts(record->start + toffset); + row.end = adjust_external_ts(record->end + toffset); row.description_id = EMPTY_STRING_ID; row.opType_id = logger.stringTable().getOrCreate("Memset"); row.api_id = record->correlationId; @@ -708,8 +708,8 @@ void CUPTIAPI CuptiDataSource::bufferCompleted(CUcontext ctx, uint32_t streamId, row.gpuId = record->deviceId; row.queueId = record->contextId; // FIXME: this or stream row.sequenceId = record->streamId; - row.start = record->start + toffset; - row.end = record->end + toffset; + row.start = adjust_external_ts(record->start + toffset); + row.end = adjust_external_ts(record->end + toffset); row.description_id = logger.stringTable().getOrCreate(cxx_demangle(record->name)); row.opType_id = logger.stringTable().getOrCreate(name); row.api_id = record->correlationId; diff --git a/rpd_tracer/DbResource.cpp b/rpd_tracer/DbResource.cpp index c97090d..641c17e 100644 --- a/rpd_tracer/DbResource.cpp +++ b/rpd_tracer/DbResource.cpp @@ -22,6 +22,7 @@ #include "DbResource.h" #include +#include "Utility.h" using rpdtracer::DbResource; @@ -47,6 +48,7 @@ DbResource::DbResource(const std::string &basefile, const std::string &resourceN : d(new DbResourcePrivate(this)) { sqlite3_open(basefile.c_str(), &d->connection); + sqlite3_busy_handler(d->connection, &rpdtracer::sqlite_busy_handler, NULL); d->resourceName = resourceName; } diff --git a/rpd_tracer/Firefly.cpp b/rpd_tracer/Firefly.cpp new file mode 100644 index 0000000..62e7c50 --- /dev/null +++ b/rpd_tracer/Firefly.cpp @@ -0,0 +1,602 @@ +#include "Firefly.h" + +#include +#include +#include + +namespace rpdtracer { +namespace firefly { + +// ----------------------------------------------------------------------------- +// Measurement Helpers +// ----------------------------------------------------------------------------- +MeasurementAnalysis read_latest_measurements(const char* role, + std::size_t windowSize, + const Measurement* measurements, + int count) { + MeasurementAnalysis analysis; + if (role == nullptr || measurements == nullptr || count <= 0) { + return analysis; + } + + analysis.samples.reserve(static_cast(count)); + for (int index = 0; index < count; ++index) { + if (std::strncmp(measurements[index].node, role, sizeof(measurements[index].node)) == 0) { + analysis.samples.push_back(measurements[index]); + } + } + + if (analysis.samples.empty()) { + return analysis; + } + + std::sort(analysis.samples.begin(), analysis.samples.end(), [](const Measurement& lhs, const Measurement& rhs) { + return lhs.timestampNs < rhs.timestampNs; + }); + + if (analysis.samples.size() > windowSize) { + analysis.samples.erase(analysis.samples.begin(), analysis.samples.end() - static_cast(windowSize)); + } + + if (analysis.samples.size() < 2U) { + analysis.averageOffset = analysis.samples.front().offset; + analysis.driftRate = 0.0; + return analysis; + } + + // Median filter: reject outlier probes before regression + { + std::vector offsets; + offsets.reserve(analysis.samples.size()); + for (const auto& s : analysis.samples) + offsets.push_back(s.offset); + std::sort(offsets.begin(), offsets.end()); + + const size_t n = offsets.size(); + const int64_t q1 = offsets[n / 4]; + const int64_t q3 = offsets[3 * n / 4]; + const int64_t iqr = q3 - q1; + const int64_t lo = q1 - 3 * iqr; + const int64_t hi = q3 + 3 * iqr; + + analysis.samples.erase( + std::remove_if(analysis.samples.begin(), analysis.samples.end(), + [lo, hi](const Measurement& m) { return m.offset < lo || m.offset > hi; }), + analysis.samples.end()); + + if (analysis.samples.size() < 2U) { + analysis.averageOffset = analysis.samples.empty() ? 0 : analysis.samples.front().offset; + analysis.driftRate = 0.0; + return analysis; + } + } + + double sumX = 0.0; + int64_t sumY = 0; + double sumXY = 0.0; + double sumX2 = 0.0; + const double referenceTime = static_cast(analysis.samples.front().timestampNs); + + for (const Measurement& measurement : analysis.samples) { + const double timeDelta = static_cast(measurement.timestampNs) - referenceTime; + int64_t offsetValue = measurement.offset; + sumX += timeDelta; + sumY += offsetValue; + sumXY += timeDelta * offsetValue; + sumX2 += timeDelta * timeDelta; + } + + const double sampleCount = static_cast(analysis.samples.size()); + const double denominator = sampleCount * sumX2 - (sumX * sumX); + if (std::fabs(denominator) < 1e-10) { + analysis.averageOffset = sumY / sampleCount; + analysis.driftRate = 0.0; + return analysis; + } + + const double slope = (sampleCount * sumXY - sumX * sumY) / denominator; + + // Use median offset instead of mean for robustness + std::vector filteredOffsets; + filteredOffsets.reserve(analysis.samples.size()); + for (const auto& s : analysis.samples) + filteredOffsets.push_back(s.offset); + std::sort(filteredOffsets.begin(), filteredOffsets.end()); + analysis.averageOffset = filteredOffsets[filteredOffsets.size() / 2]; + analysis.driftRate = slope; + return analysis; +} + +// ----------------------------------------------------------------------------- +// Networking Utilities +// ----------------------------------------------------------------------------- +bool enable_sw_timestamps(int sockFd) { + const int timestampOptions = SOF_TIMESTAMPING_TX_SOFTWARE | + SOF_TIMESTAMPING_RX_SOFTWARE | + SOF_TIMESTAMPING_SOFTWARE | + SOF_TIMESTAMPING_OPT_TSONLY; + + if (setsockopt(sockFd, SOL_SOCKET, SO_TIMESTAMPING, ×tampOptions, sizeof(timestampOptions)) < 0) { +// std::fprintf(stderr, "ChronoSync: [enable_sw_timestamps] CRITICAL - setsockopt failed (errno: %s)\n", std::strerror(errno)); + return false; + } + return true; +} + +bool send_probe(int sockFd, sockaddr_in* peerAddress, timestamp_t* sendTime, int probeId) { + if (peerAddress == nullptr || sendTime == nullptr) { +// std::fprintf(stderr, "ChronoSync: [send_probe] CRITICAL - Invalid arguments\n"); + return false; + } + + char messageBuffer[NETWORK_BUFFER_SIZE] = {}; + if (std::snprintf(messageBuffer, sizeof(messageBuffer), "Probe %d", probeId) < 0) { +// std::fprintf(stderr, "ChronoSync: [send_probe] CRITICAL - snprintf failed\n"); + return false; + } + + if (sendto(sockFd, + messageBuffer, + std::strlen(messageBuffer), + 0, + reinterpret_cast(peerAddress), + sizeof(*peerAddress)) < 0) { +// std::fprintf(stderr, "ChronoSync: [send_probe] CRITICAL - sendto failed (errno: %s)\n", std::strerror(errno)); + return false; + } + + timespec timeSpec{}; + if (clock_gettime(CLOCK_MONOTONIC, &timeSpec) == 0) { + *sendTime = timespec_to_ns(timeSpec); + } else { +// std::fprintf(stderr, "ChronoSync: [send_probe] WARNING - clock_gettime failed (errno: %s)\n", std::strerror(errno)); + *sendTime = 0; + } + return true; +} + +void receive_probe(int sockFd, timestamp_t* receiveTime, int expectedProbeId) { + if (receiveTime == nullptr) { +// std::fprintf(stderr, "ChronoSync: [receive_probe] CRITICAL - Invalid receiveTime pointer\n"); + return; + } + + char messageBuffer[NETWORK_BUFFER_SIZE] = {}; + sockaddr_in senderAddress{}; + msghdr messageHeader{}; + iovec bufferVector[1]; + char controlBuffer[NETWORK_BUFFER_SIZE] = {}; + + bufferVector[0].iov_base = messageBuffer; + bufferVector[0].iov_len = sizeof(messageBuffer); + + messageHeader.msg_iov = bufferVector; + messageHeader.msg_iovlen = 1; + messageHeader.msg_name = &senderAddress; + messageHeader.msg_namelen = sizeof(senderAddress); + messageHeader.msg_control = controlBuffer; + messageHeader.msg_controllen = sizeof(controlBuffer); + + const ssize_t receivedBytes = recvmsg(sockFd, &messageHeader, 0); + if (receivedBytes < 0) { + *receiveTime = 0; + return; + } + + if (static_cast(receivedBytes) >= sizeof(messageBuffer)) { + *receiveTime = 0; + return; + } + + messageBuffer[receivedBytes] = '\0'; + int receivedProbeId = 0; + if (std::sscanf(messageBuffer, "Probe %d", &receivedProbeId) != 1 || receivedProbeId != expectedProbeId) { + *receiveTime = 0; + return; + } + + timespec timeSpec{}; + if (clock_gettime(CLOCK_MONOTONIC, &timeSpec) == 0) { + *receiveTime = timespec_to_ns(timeSpec); + } else { +// std::fprintf(stderr, "ChronoSync: [receive_probe] WARNING - clock_gettime failed (errno: %s)\n", std::strerror(errno)); + *receiveTime = 0; + } +} + +int tcp_handshake(const char* role, const char* peerIp, int tcpPort) { + if (role == nullptr || peerIp == nullptr) + return -1; + + std::fprintf(stderr, "ChronoSync: connecting to %s (role %s, port %d)\n", peerIp, role, tcpPort); + const bool isNodeA = (std::strcmp(role, "A") == 0); + int tcpSocketFd = -1; + sockaddr_in tcpAddress{}; + tcpAddress.sin_family = AF_INET; + tcpAddress.sin_port = htons(tcpPort); + + if (isNodeA) { + tcpSocketFd = socket(AF_INET, SOCK_STREAM, 0); + if (tcpSocketFd < 0) { +// std::fprintf(stderr, "ChronoSync: [tcp_handshake] CRITICAL - socket failed (errno: %s)\n", std::strerror(errno)); + return -1; + } + + if (inet_aton(peerIp, &tcpAddress.sin_addr) == 0) { +// std::fprintf(stderr, "ChronoSync: [tcp_handshake] CRITICAL - Invalid peer IP\n"); + close(tcpSocketFd); + return -1; + } + + bool connected = false; + for (int attempt = 0; attempt < CONNECTION_RETRY_LIMIT; ++attempt) { + if (connect(tcpSocketFd, reinterpret_cast(&tcpAddress), sizeof(tcpAddress)) == 0) { + connected = true; + break; + } +// std::fprintf(stderr, "ChronoSync: [tcp_handshake] WARNING - connect attempt %d failed (errno: %s)\n", +// attempt + 1, +// std::strerror(errno)); + sleep(CONNECTION_RETRY_DELAY_SEC); + } + + if (!connected) { + std::fprintf(stderr, "ChronoSync: failed to connect to %s:%d after %d attempts\n", peerIp, tcpPort, CONNECTION_RETRY_LIMIT); + close(tcpSocketFd); + return -1; + } + + constexpr char HANDSHAKE_READY[] = "READY"; + if (send(tcpSocketFd, HANDSHAKE_READY, sizeof(HANDSHAKE_READY), 0) < 0) { +// std::fprintf(stderr, "ChronoSync: [tcp_handshake] CRITICAL - send failed (errno: %s)\n", std::strerror(errno)); + close(tcpSocketFd); + return -1; + } + + char ackBuffer[NETWORK_BUFFER_SIZE] = {}; + const ssize_t received = recv(tcpSocketFd, ackBuffer, sizeof(ackBuffer) - 1, 0); + if (received <= 0) { +// std::fprintf(stderr, "ChronoSync: [tcp_handshake] WARNING - recv failed (errno: %s)\n", std::strerror(errno)); + } else { + ackBuffer[received] = '\0'; + } + } else { + int listenSocketFd = socket(AF_INET, SOCK_STREAM, 0); + if (listenSocketFd < 0) { +// std::fprintf(stderr, "ChronoSync: [tcp_handshake] CRITICAL - socket failed (errno: %s)\n", std::strerror(errno)); + return -1; + } + + int reuseAddress = 1; + if (setsockopt(listenSocketFd, SOL_SOCKET, SO_REUSEADDR, &reuseAddress, sizeof(reuseAddress)) < 0) { +// std::fprintf(stderr, "ChronoSync: [tcp_handshake] CRITICAL - setsockopt failed (errno: %s)\n", std::strerror(errno)); + close(listenSocketFd); + return -1; + } + + tcpAddress.sin_addr.s_addr = INADDR_ANY; + if (bind(listenSocketFd, reinterpret_cast(&tcpAddress), sizeof(tcpAddress)) < 0) { +// std::fprintf(stderr, "ChronoSync: [tcp_handshake] CRITICAL - bind failed (errno: %s)\n", std::strerror(errno)); + close(listenSocketFd); + return -1; + } + + if (listen(listenSocketFd, 1) < 0) { +// std::fprintf(stderr, "ChronoSync: [tcp_handshake] CRITICAL - listen failed (errno: %s)\n", std::strerror(errno)); + close(listenSocketFd); + return -1; + } + + sockaddr_in clientAddress{}; + socklen_t clientLength = sizeof(clientAddress); + tcpSocketFd = accept(listenSocketFd, reinterpret_cast(&clientAddress), &clientLength); + close(listenSocketFd); + + if (tcpSocketFd < 0) { +// std::fprintf(stderr, "ChronoSync: [tcp_handshake] CRITICAL - accept failed (errno: %s)\n", std::strerror(errno)); + return -1; + } + + char handshakeBuffer[NETWORK_BUFFER_SIZE] = {}; + const ssize_t received = recv(tcpSocketFd, handshakeBuffer, sizeof(handshakeBuffer) - 1, 0); + if (received > 0) { + handshakeBuffer[received] = '\0'; + } + + constexpr char HANDSHAKE_ACK[] = "ACK"; + if (send(tcpSocketFd, HANDSHAKE_ACK, sizeof(HANDSHAKE_ACK), 0) < 0) { +// std::fprintf(stderr, "ChronoSync: [tcp_handshake] CRITICAL - send failed (errno: %s)\n", std::strerror(errno)); + close(tcpSocketFd); + return -1; + } + } + + std::fprintf(stderr, "ChronoSync: connected to %s (role %s, port %d)\n", peerIp, role, tcpPort); + return tcpSocketFd; +} + +// ----------------------------------------------------------------------------- +// Clock Helpers +// ----------------------------------------------------------------------------- + +timespec hw_now() { + timespec timeSpec{}; + if (clock_gettime(CLOCK_MONOTONIC, &timeSpec) != 0) { +// std::fprintf(stderr, "ChronoSync: [hw_now] WARNING - clock_gettime failed (errno: %s)\n", std::strerror(errno)); + } + return timeSpec; +} + +void svc_update_ns(SvcState* state, int64_t offsetNs, double drift) { + if (state == nullptr) { +// std::fprintf(stderr, "ChronoSync: [svc_update_ns] CRITICAL - Null svc state\n"); + return; + } + + const unsigned sequence = state->sequence.fetch_add(1U, std::memory_order_acq_rel); + state->referenceTime = hw_now(); + state->offset = offsetNs; + state->drift = drift; + state->sequence.store(sequence + 2U, std::memory_order_release); +} + +// ----------------------------------------------------------------------------- +// Firefly +// ----------------------------------------------------------------------------- +void firefly_run(const char* role, MeasurementBuffer& buffer) { + if (role == nullptr) { +// std::fprintf(stderr, "ChronoSync: [firefly_run] CRITICAL - Invalid arguments\n"); + return; + } + + int localCount = 0; + { + std::lock_guard lock(buffer.mutex); + localCount = buffer.count; + } + + const MeasurementAnalysis analysis = read_latest_measurements(role, + REGRESSION_WINDOW_SIZE, + buffer.measurements.data(), + localCount); + + // Real clock drift is < 100 ppm; anything larger indicates bad data + double drift = analysis.driftRate; + if (std::fabs(drift) > 500e-6) { +// std::fprintf(stderr, "ChronoSync: [firefly_run] WARNING - Drift %.9e exceeds 500 ppm, clamping to 0\n", drift); + drift = 0.0; + } + + svc_update_ns(firefly::g_pSvcState, static_cast(analysis.averageOffset * CONSENSUS_ALPHA), drift * CONSENSUS_ALPHA); +// std::fprintf(stderr, +// "ChronoSync: [firefly_run] INFO - Role %s, samples=%zu, offset=%lld, drift=%.9e\n", +// role, +// analysis.samples.size(), +// static_cast(firefly::g_pSvcState->offset), +// firefly::g_pSvcState->drift); +} + +// ----------------------------------------------------------------------------- +// UDP Node +// ----------------------------------------------------------------------------- +static void record_measurement(MeasurementBuffer& buffer, const Measurement& measurement) { + std::lock_guard lock(buffer.mutex); + if (buffer.count < static_cast(buffer.measurements.size())) { + buffer.measurements[buffer.count] = measurement; + ++buffer.count; + } +} + +void run_node(const char* role, + const char* ipA, + const char* ipB, + int udpPort, + MeasurementBuffer& buffer, + bool& done) { + if (role == nullptr || ipA == nullptr || ipB == nullptr) { +// std::fprintf(stderr, "ChronoSync: [run_node] CRITICAL - Invalid arguments\n"); + return; + } + + const bool isNodeA = (std::strcmp(role, "A") == 0); + int socketFd = socket(AF_INET, SOCK_DGRAM, 0); + if (socketFd < 0) { +// std::fprintf(stderr, "ChronoSync: [run_node] CRITICAL - socket failed (errno: %s)\n", std::strerror(errno)); + return; + } + + timeval timeoutValue{}; + timeoutValue.tv_sec = SOCKET_TIMEOUT_MSEC / 1000; + timeoutValue.tv_usec = (SOCKET_TIMEOUT_MSEC % 1000) * 1000; + if (setsockopt(socketFd, SOL_SOCKET, SO_RCVTIMEO, &timeoutValue, sizeof(timeoutValue)) < 0) { +// std::fprintf(stderr, "ChronoSync: [run_node] CRITICAL - setsockopt failed (errno: %s)\n", std::strerror(errno)); + close(socketFd); + return; + } + + if (!enable_sw_timestamps(socketFd)) { + close(socketFd); + return; + } + + sockaddr_in localAddress{}; + sockaddr_in peerAddress{}; + localAddress.sin_family = AF_INET; + peerAddress.sin_family = AF_INET; + localAddress.sin_port = htons(udpPort); + peerAddress.sin_port = htons(udpPort); + + const char* localIp = isNodeA ? ipA : ipB; + const char* peerIp = isNodeA ? ipB : ipA; + + if (inet_aton(localIp, &localAddress.sin_addr) == 0 || inet_aton(peerIp, &peerAddress.sin_addr) == 0) { +// std::fprintf(stderr, "ChronoSync: [run_node] CRITICAL - Invalid IP address\n"); + close(socketFd); + return; + } + + if (bind(socketFd, reinterpret_cast(&localAddress), sizeof(localAddress)) < 0) { +// std::fprintf(stderr, "ChronoSync: [run_node] CRITICAL - bind failed (errno: %s)\n", std::strerror(errno)); + close(socketFd); + return; + } + + int tcpSocketFd = tcp_handshake(role, peerIp, udpPort); + if (tcpSocketFd < 0) { + close(socketFd); + return; + } + close(tcpSocketFd); + + char messageBuffer[NETWORK_BUFFER_SIZE]; + int skippedIterations = 0; + int probeIndex = 0; + + while (done == false) { + ++probeIndex; + timestamp_t sendTimeA = 0; + timestamp_t recvTimeA = 0; + timestamp_t sendTimeB = 0; + timestamp_t recvTimeB = 0; + + if (isNodeA) { + if (!send_probe(socketFd, &peerAddress, &sendTimeA, probeIndex)) { + ++skippedIterations; + continue; + } + receive_probe(socketFd, &recvTimeA, probeIndex); + if (recvTimeA == 0) { + ++skippedIterations; + continue; + } + + if (std::snprintf(messageBuffer, sizeof(messageBuffer), "%lld %lld", + static_cast(sendTimeA), + static_cast(recvTimeA)) < 0) { + continue; + } + + if (sendto(socketFd, messageBuffer, std::strlen(messageBuffer), 0, + reinterpret_cast(&peerAddress), sizeof(peerAddress)) < 0) { + ++skippedIterations; + continue; + } + + sockaddr_in senderAddress{}; + socklen_t senderLength = sizeof(senderAddress); + const ssize_t receivedBytes = recvfrom(socketFd, + messageBuffer, + sizeof(messageBuffer) - 1, + 0, + reinterpret_cast(&senderAddress), + &senderLength); + if (receivedBytes < 0) { + ++skippedIterations; + continue; + } + + messageBuffer[receivedBytes] = '\0'; + if (std::sscanf(messageBuffer, "%lld %lld", + reinterpret_cast(&sendTimeB), + reinterpret_cast(&recvTimeB)) != 2) { + ++skippedIterations; + continue; + } + + const timestamp_t roundTripTime = (recvTimeA - sendTimeA) - (sendTimeB - recvTimeB); + const int64_t offset = (recvTimeB - sendTimeA) - (roundTripTime / 2); + + timespec now{}; + clock_gettime(CLOCK_MONOTONIC, &now); + Measurement measurement{}; + measurement.timestampNs = timespec_to_ns(now); + measurement.sendTimeA = sendTimeA; + measurement.recvTimeA = recvTimeA; + measurement.sendTimeB = sendTimeB; + measurement.recvTimeB = recvTimeB; + measurement.roundTripTime = roundTripTime; + measurement.offset = offset; + measurement.udpPort = udpPort; + measurement.node[0] = 'A'; + measurement.node[1] = '\0'; + + record_measurement(buffer, measurement); + } else { + receive_probe(socketFd, &recvTimeB, probeIndex); + if (recvTimeB == 0) { + ++skippedIterations; + continue; + } + + if (!send_probe(socketFd, &peerAddress, &sendTimeB, probeIndex)) { + ++skippedIterations; + continue; + } + + sockaddr_in senderAddress{}; + socklen_t senderLength = sizeof(senderAddress); + const ssize_t receivedBytes = recvfrom(socketFd, + messageBuffer, + sizeof(messageBuffer) - 1, + 0, + reinterpret_cast(&senderAddress), + &senderLength); + if (receivedBytes < 0) { + ++skippedIterations; + continue; + } + + messageBuffer[receivedBytes] = '\0'; + if (std::sscanf(messageBuffer, "%lld %lld", + reinterpret_cast(&sendTimeA), + reinterpret_cast(&recvTimeA)) != 2) { + ++skippedIterations; + continue; + } + + if (std::snprintf(messageBuffer, sizeof(messageBuffer), "%lld %lld", + static_cast(sendTimeB), + static_cast(recvTimeB)) < 0) { + ++skippedIterations; + continue; + } + + if (sendto(socketFd, + messageBuffer, + std::strlen(messageBuffer), + 0, + reinterpret_cast(&peerAddress), + sizeof(peerAddress)) < 0) { + ++skippedIterations; + continue; + } + + const timestamp_t roundTripTime = (recvTimeA - sendTimeA) - (sendTimeB - recvTimeB); + const int64_t offset = (recvTimeA - sendTimeB) - (roundTripTime / 2); + + timespec now{}; + clock_gettime(CLOCK_MONOTONIC, &now); + Measurement measurement{}; + measurement.timestampNs = timespec_to_ns(now); + measurement.sendTimeA = sendTimeA; + measurement.recvTimeA = recvTimeA; + measurement.sendTimeB = sendTimeB; + measurement.recvTimeB = recvTimeB; + measurement.roundTripTime = roundTripTime; + measurement.offset = offset; + measurement.udpPort = udpPort; + measurement.node[0] = 'B'; + measurement.node[1] = '\0'; + + record_measurement(buffer, measurement); + } + + std::this_thread::sleep_for(std::chrono::milliseconds(100)); + } + + close(socketFd); +// std::fprintf(stderr, "ChronoSync: [run_node] INFO - Node %s completed, skipped %d iterations\n", role, skippedIterations); +} + +} // namespace firefly +} // namespace rpdtracer diff --git a/rpd_tracer/Firefly.h b/rpd_tracer/Firefly.h new file mode 100644 index 0000000..dca62d3 --- /dev/null +++ b/rpd_tracer/Firefly.h @@ -0,0 +1,99 @@ +#ifndef FIREFLY_H +#define FIREFLY_H + +#include +#include +#include +#include +#include +#include +#include + +#include +#include +#include +#include +#include +#include +#include +#include "Utility.h" + +struct sockaddr_in; + +namespace rpdtracer { +namespace firefly { + +using timestamp_t = uint64_t; + +// FIXME: use rlog properties in their own domain +// FIXME: replace batch regression with a PID controller for smooth, continuous +// clock correction. The PID naturally slews (no offset discontinuities) and +// eliminates the need for windowed regression entirely. +constexpr int TCP_PORT_DEFAULT = 12345; +constexpr int UDP_PORT_DEFAULT = 12345; +constexpr std::size_t NETWORK_BUFFER_SIZE = 1024U; +constexpr int SOCKET_TIMEOUT_MSEC = 200; +constexpr int CONNECTION_RETRY_LIMIT = 30; +constexpr int CONNECTION_RETRY_DELAY_SEC = 1; +constexpr std::size_t MAX_MEASUREMENT_COUNT = 100000U; +constexpr std::size_t REGRESSION_WINDOW_SIZE = 200U; +constexpr double CONSENSUS_ALPHA = 0.5; +constexpr double GAIN_PHASE = 1.0; +constexpr double GAIN_FREQ = 1.0; +constexpr int FIRE_FLY_SLEEP_MSEC = 100; + +struct Measurement { + timestamp_t timestampNs; + timestamp_t sendTimeA; + timestamp_t recvTimeA; + timestamp_t sendTimeB; + timestamp_t recvTimeB; + timestamp_t roundTripTime; + int64_t offset; + int udpPort; + char node[2]; +}; + +struct MeasurementAnalysis { + std::vector samples; + int64_t averageOffset{0}; + double driftRate{0.0}; +}; + +struct MeasurementBuffer { + std::vector measurements; + int count{0}; + std::mutex mutex; + + MeasurementBuffer(std::size_t capacity) : measurements(capacity) {} +}; + +bool enable_sw_timestamps(int sockFd); + +bool send_probe(int sockFd, sockaddr_in* peerAddress, timestamp_t* sendTime, int probeId); + +void receive_probe(int sockFd, timestamp_t* receiveTime, int expectedProbeId); + +int tcp_handshake(const char* role, const char* peerIp, int tcpPort); + +void run_node(const char* role, + const char* ipA, + const char* ipB, + int udpPort, + MeasurementBuffer& buffer, + bool& done); + +void firefly_run(const char* role, MeasurementBuffer& buffer); + +timespec hw_now(); + +void svc_update_ns(SvcState* state, int64_t offsetNs, double drift); + +MeasurementAnalysis read_latest_measurements(const char* role, + std::size_t windowSize, + const Measurement* measurements, + int count); +} // namespace firefly +} // namespace rpdtracer + +#endif // FIREFLY_H diff --git a/rpd_tracer/Logger.cpp b/rpd_tracer/Logger.cpp index caf0b42..fc15ec0 100644 --- a/rpd_tracer/Logger.cpp +++ b/rpd_tracer/Logger.cpp @@ -249,6 +249,10 @@ void Logger::init() "RocmSmiDataSourceFactory" }; + // FIXME: use rlog property + if (getenv("RPDT_CLOCKSYNC_IP") != nullptr) + factories.push_back("ChronoSyncDataSourceFactory"); + for (auto it = factories.begin(); it != factories.end(); ++it) { DataSource* (*func) (void) = (DataSource* (*)()) dlsym(RTLD_DEFAULT, (*it).c_str()); if (func) { diff --git a/rpd_tracer/Logger.h b/rpd_tracer/Logger.h index cd33326..fb9b917 100644 --- a/rpd_tracer/Logger.h +++ b/rpd_tracer/Logger.h @@ -37,6 +37,10 @@ const sqlite_int64 EMPTY_STRING_ID = 1; class Logger { public: + // FIXME: calling init() from the constructor deadlocks any data source that + // calls Logger::singleton() during its own init(), because the static init + // guard is still held. Fix: empty constructor, move init() into singleton() + // after the static, use m_initialized flag to make init() idempotent. Logger() { init(); } static Logger& singleton(); diff --git a/rpd_tracer/Makefile b/rpd_tracer/Makefile index 59c0688..2b7bc71 100644 --- a/rpd_tracer/Makefile +++ b/rpd_tracer/Makefile @@ -9,7 +9,7 @@ ROCM_VERSION ?= $(shell cat /opt/rocm/.info/version) RPD_LIBS = -lsqlite3 -lfmt RPD_INCLUDES = -RPD_SRCS = Schema.cpp Table.cpp BufferedTable.cpp OpTable.cpp KernelApiTable.cpp CopyApiTable.cpp ApiTable.cpp StringTable.cpp UStringTable.cpp MetadataTable.cpp MonitorTable.cpp StackFrameTable.cpp ApiIdList.cpp DbResource.cpp Logger.cpp Unwind.cpp RoctxDataSource.cpp NvtxDataSource.cpp +RPD_SRCS = Schema.cpp Table.cpp BufferedTable.cpp OpTable.cpp KernelApiTable.cpp CopyApiTable.cpp ApiTable.cpp StringTable.cpp UStringTable.cpp MetadataTable.cpp MonitorTable.cpp StackFrameTable.cpp ApiIdList.cpp DbResource.cpp Logger.cpp Unwind.cpp RoctxDataSource.cpp NvtxDataSource.cpp ChronoSyncDataSource.cpp Firefly.cpp Utility.cpp ifneq (,$(HIP_PATH)) $(info rocm_version is ${ROCM_VERSION}) diff --git a/rpd_tracer/RocmSmiDataSource.cpp b/rpd_tracer/RocmSmiDataSource.cpp index 6824943..00abb30 100644 --- a/rpd_tracer/RocmSmiDataSource.cpp +++ b/rpd_tracer/RocmSmiDataSource.cpp @@ -43,6 +43,7 @@ void RocmSmiDataSource::init() m_done = false; m_period = 1000; + // FIXME: Logger::singleton() deadlocks here — called during init() while singleton is constructing m_resource = new DbResource(Logger::singleton().filename(), std::string("smi_logger_active")); m_worker = new std::thread(&RocmSmiDataSource::work, this); } diff --git a/rpd_tracer/RocprofDataSource.cpp b/rpd_tracer/RocprofDataSource.cpp index 8967b26..5cbbb64 100644 --- a/rpd_tracer/RocprofDataSource.cpp +++ b/rpd_tracer/RocprofDataSource.cpp @@ -503,8 +503,8 @@ void RocprofDataSource::api_callback(rocprofiler_callback_tracing_record_t recor row.queueId = info.queue_id.handle; row.sequenceId = info.dispatch_id; strncpy(row.completionSignal, "", 18); - row.start = dispatch.start_timestamp; - row.end = dispatch.end_timestamp; + row.start = adjust_external_ts(dispatch.start_timestamp); + row.end = adjust_external_ts(dispatch.end_timestamp); row.description_id = logger.stringTable().getOrCreate(s->kernel_names.at(info.kernel_id)); row.opType_id = name_id; row.api_id = record.correlation_id.internal; @@ -543,8 +543,8 @@ void RocprofDataSource::api_callback(rocprofiler_callback_tracing_record_t recor row.queueId = 0; row.sequenceId = 0; strncpy(row.completionSignal, "", 18); - row.start = copy.start_timestamp; - row.end = copy.end_timestamp; + row.start = adjust_external_ts(copy.start_timestamp); + row.end = adjust_external_ts(copy.end_timestamp); row.description_id = logger.stringTable().getOrCreate(crow.kindStr); row.opType_id = name_id; row.api_id = record.correlation_id.internal; @@ -588,8 +588,8 @@ void RocprofDataSource::buffer_callback(rocprofiler_context_id_t context, rocpro row.gpuId = s->agents.at(dispatch.agent_id.handle).logical_node_type_id; row.queueId = dispatch.queue_id.handle; row.sequenceId = 0; - row.start = record->start_timestamp; - row.end = record->end_timestamp; + row.start = adjust_external_ts(record->start_timestamp); + row.end = adjust_external_ts(record->end_timestamp); row.description_id = desc_id; row.opType_id = name_id; row.api_id = record->correlation_id.internal; @@ -627,8 +627,8 @@ void RocprofDataSource::buffer_callback(rocprofiler_context_id_t context, rocpro row.gpuId = 0; row.queueId = 0; // FIXME, all wrong row.sequenceId = 0; - row.start = copy.start_timestamp; - row.end = copy.end_timestamp; + row.start = adjust_external_ts(copy.start_timestamp); + row.end = adjust_external_ts(copy.end_timestamp); row.description_id = desc_id; row.opType_id = name_id; row.api_id = copy.correlation_id.internal; @@ -671,8 +671,8 @@ void RocprofDataSource::buffer_callback(rocprofiler_context_id_t context, rocpro ApiTable::row row; row.pid = GetPid(); row.tid = hipapi.thread_id; - row.start = hipapi.start_timestamp; - row.end = hipapi.end_timestamp; + row.start = adjust_external_ts(hipapi.start_timestamp); + row.end = adjust_external_ts(hipapi.end_timestamp); row.domain_id = domain_id; row.category_id = EMPTY_STRING_ID; row.apiName_id = name_id; diff --git a/rpd_tracer/RoctracerDataSource.cpp b/rpd_tracer/RoctracerDataSource.cpp index 979a2ab..08b5287 100644 --- a/rpd_tracer/RoctracerDataSource.cpp +++ b/rpd_tracer/RoctracerDataSource.cpp @@ -795,8 +795,8 @@ return; sqlite3_bind_int(apiInsert, index++, record->correlation_id); sqlite3_bind_int(apiInsert, index++, record->process_id); sqlite3_bind_int(apiInsert, index++, record->thread_id); - sqlite3_bind_int64(apiInsert, index++, record->begin_ns); - sqlite3_bind_int64(apiInsert, index++, record->end_ns); + sqlite3_bind_int64(apiInsert, index++, adjust_external_ts(record->begin_ns)); + sqlite3_bind_int64(apiInsert, index++, adjust_external_ts(record->end_ns)); sqlite3_bind_int64(apiInsert, index++, rowId); sqlite3_bind_int64(apiInsert, index++, EMPTY_STRING_ID); @@ -829,16 +829,6 @@ void RoctracerDataSource::hcc_activity_callback(const char* begin, const char* e int batchSize = 0; - // Roctracer uses CLOCK_MONOTONIC for timestamps, which matches our timestamps. - // However, the Roctracer developer thinks it is using CLOCK_MONOTONIC_RAW, which it isn't - // Go ahead and convert timestamps here just in case this gets "fixed" at some point - timestamp_t t0, t1, t00; - roctracer_get_timestamp(&t1); // first call is really slow, throw it away - t0 = clocktime_ns(); - roctracer_get_timestamp(&t1); - t00 = clocktime_ns(); - const timestamp_t toffset = (t0 >> 1) + (t00 >> 1) - t1; - Logger &logger = Logger::singleton(); while (record < end_record) { @@ -850,8 +840,8 @@ void RoctracerDataSource::hcc_activity_callback(const char* begin, const char* e row.gpuId = mapDeviceId(record->device_id); row.queueId = record->queue_id; row.sequenceId = 0; - row.start = record->begin_ns + toffset; - row.end = record->end_ns + toffset; + row.start = adjust_external_ts(record->begin_ns); + row.end = adjust_external_ts(record->end_ns); row.description_id = ((record->kind == HIP_OP_DISPATCH_KIND_KERNEL_) || (record->kind == HIP_OP_DISPATCH_KIND_TASK_)) ? logger.stringTable().getOrCreate(cxx_demangle(record->kernel_name)) diff --git a/rpd_tracer/Schema.cpp b/rpd_tracer/Schema.cpp index da045b6..0ecd38e 100644 --- a/rpd_tracer/Schema.cpp +++ b/rpd_tracer/Schema.cpp @@ -7,6 +7,7 @@ #include #include +#include "Utility.h" #include "tableSchema.h" #include "utilitySchema.h" @@ -16,6 +17,8 @@ void ensureSchema(const char *basefile) { sqlite3 *db = nullptr; int ret = sqlite3_open(basefile, &db); + if (ret == SQLITE_OK) + sqlite3_busy_handler(db, &sqlite_busy_handler, NULL); if (ret != SQLITE_OK) { fprintf(stderr, "rpd_tracer: cannot open database %s: %s\n", basefile, sqlite3_errmsg(db)); return; diff --git a/rpd_tracer/Table.cpp b/rpd_tracer/Table.cpp index 0bf0395..e4c3ca5 100644 --- a/rpd_tracer/Table.cpp +++ b/rpd_tracer/Table.cpp @@ -7,12 +7,6 @@ using rpdtracer::Table; -int busy_handler(void *data, int count) -{ - count = (count < 9) ? count : 8; - usleep(1000 * (0x1 << count)); - return 1; -} static int wal_check_callback(void *data, int ncols, char **values, char **names) { @@ -25,7 +19,7 @@ Table::Table(const char *basefile) : m_connection(NULL) { sqlite3_open(basefile, &m_connection); - sqlite3_busy_handler(m_connection, &busy_handler, NULL); + sqlite3_busy_handler(m_connection, &sqlite_busy_handler, NULL); bool walEnabled = false; sqlite3_exec(m_connection, "PRAGMA journal_mode=WAL", wal_check_callback, &walEnabled, NULL); diff --git a/rpd_tracer/Utility.cpp b/rpd_tracer/Utility.cpp new file mode 100644 index 0000000..5da3a8b --- /dev/null +++ b/rpd_tracer/Utility.cpp @@ -0,0 +1,101 @@ +/************************************************************************** + * Copyright (c) 2023 Advanced Micro Devices, Inc. + **************************************************************************/ + +#include "Utility.h" + +#include +#include +#include +#include +#include + +namespace rpdtracer { + +int sqlite_busy_handler(void *data, int count) +{ + count = (count < 9) ? count : 8; + usleep(1000 * (0x1 << count)); + return 1; +} + +namespace firefly { + +static SvcState s_localState; +SvcState* g_pSvcState = &s_localState; + +static void* s_shmAddr = nullptr; + +void create_clocksync_shm(std::string& shm_name_out) { + timespec ts; + clock_gettime(CLOCK_MONOTONIC, &ts); + char name[64]; + std::snprintf(name, sizeof(name), "/rpd_sync_%d_%lld", + static_cast(getpid()), + static_cast(ts.tv_sec * 1000000000LL + ts.tv_nsec)); + shm_name_out = name; + + int fd = shm_open(name, O_CREAT | O_EXCL | O_RDWR, 0666); + if (fd < 0) { +// std::fprintf(stderr, "ChronoSync: shm_open create failed (errno: %s)\n", std::strerror(errno)); + return; + } + + if (ftruncate(fd, sizeof(SvcState)) < 0) { +// std::fprintf(stderr, "ChronoSync: ftruncate failed (errno: %s)\n", std::strerror(errno)); + close(fd); + shm_unlink(name); + return; + } + + void* addr = mmap(nullptr, sizeof(SvcState), PROT_READ | PROT_WRITE, MAP_SHARED, fd, 0); + close(fd); + if (addr == MAP_FAILED) { +// std::fprintf(stderr, "ChronoSync: mmap failed (errno: %s)\n", std::strerror(errno)); + shm_unlink(name); + return; + } + + std::memset(addr, 0, sizeof(SvcState)); + new (addr) SvcState(); + + s_shmAddr = addr; + g_pSvcState = static_cast(addr); +// std::fprintf(stderr, "ChronoSync: created shm %s\n", name); +} + +void attach_clocksync_shm(const std::string& shm_name) { + int fd = shm_open(shm_name.c_str(), O_RDWR, 0); + if (fd < 0) { +// std::fprintf(stderr, "ChronoSync: shm_open attach failed for %s (errno: %s)\n", +// shm_name.c_str(), std::strerror(errno)); + return; + } + + void* addr = mmap(nullptr, sizeof(SvcState), PROT_READ | PROT_WRITE, MAP_SHARED, fd, 0); + close(fd); + if (addr == MAP_FAILED) { +// std::fprintf(stderr, "ChronoSync: mmap attach failed (errno: %s)\n", std::strerror(errno)); + return; + } + + s_shmAddr = addr; + g_pSvcState = static_cast(addr); +// std::fprintf(stderr, "ChronoSync: attached shm %s\n", shm_name.c_str()); +} + +void cleanup_clocksync_shm(const std::string& shm_name) { + if (s_shmAddr != nullptr) { + munmap(s_shmAddr, sizeof(SvcState)); + s_shmAddr = nullptr; + } + g_pSvcState = &s_localState; + + if (!shm_name.empty()) + shm_unlink(shm_name.c_str()); + +// std::fprintf(stderr, "ChronoSync: cleaned up shm %s\n", shm_name.c_str()); +} + +} // namespace firefly +} // namespace rpdtracer diff --git a/rpd_tracer/Utility.h b/rpd_tracer/Utility.h index 8b7e34c..95fc7a3 100644 --- a/rpd_tracer/Utility.h +++ b/rpd_tracer/Utility.h @@ -9,6 +9,8 @@ #include #include #include +#include +#include #include namespace rpdtracer { @@ -39,12 +41,58 @@ static timestamp_t timespec_to_ns(const timespec& time) { return ((timestamp_t)time.tv_sec * 1000000000) + time.tv_nsec; } -static timestamp_t clocktime_ns() { +// Seqlock-protected clock sync state, written by ChronoSync's Firefly algorithm. +// Inlined here so the seq==0 fast path (no sync active) costs only a single +// load+branch on the ~80 clocktime_ns() call sites in the tracing hot path. +namespace firefly { + +struct SvcState { + std::atomic sequence{0}; + timespec referenceTime{}; + int64_t offset{0}; + double drift{0.0}; +}; + +extern SvcState* g_pSvcState; + +void create_clocksync_shm(std::string& shm_name_out); +void attach_clocksync_shm(const std::string& shm_name); +void cleanup_clocksync_shm(const std::string& shm_name); + +} // namespace firefly + +static inline int64_t svc_read_offset_ns(timestamp_t now_ns) { + unsigned seq = firefly::g_pSvcState->sequence.load(std::memory_order_acquire); + if (seq == 0) + return 0; + + int64_t offset; + double drift; + timestamp_t ref_ns; + do { + seq = firefly::g_pSvcState->sequence.load(std::memory_order_acquire); + offset = firefly::g_pSvcState->offset; + drift = firefly::g_pSvcState->drift; + ref_ns = timespec_to_ns(firefly::g_pSvcState->referenceTime); + } while ((seq & 1) || firefly::g_pSvcState->sequence.load(std::memory_order_acquire) != seq); + + int64_t elapsed = static_cast(now_ns - ref_ns); + return offset + static_cast(drift * elapsed); +} + +static inline timestamp_t adjust_external_ts(timestamp_t ts) { + return ts + svc_read_offset_ns(ts); +} + +static inline timestamp_t clocktime_ns() { timespec ts; clock_gettime(CLOCK_MONOTONIC, &ts); - return ((timestamp_t)ts.tv_sec * 1000000000) + ts.tv_nsec; + timestamp_t now_ns = ((timestamp_t)ts.tv_sec * 1000000000) + ts.tv_nsec; + return now_ns + svc_read_offset_ns(now_ns); } +int sqlite_busy_handler(void *data, int count); + void createOverheadRecord(uint64_t start, uint64_t end, const std::string &name, const std::string &args); class Logger; diff --git a/rpd_tracer/tests/Makefile b/rpd_tracer/tests/Makefile new file mode 100644 index 0000000..7ac64b5 --- /dev/null +++ b/rpd_tracer/tests/Makefile @@ -0,0 +1,35 @@ +CXX = /opt/rocm/llvm/bin/clang++ +CXXFLAGS = -std=c++11 -g -O2 -I.. -pthread + +TESTS = test_seqlock test_measurement test_regression + +# Firefly.o and Utility.o built from parent directory +PARENT_OBJS = ../Firefly.o ../Utility.o + +.PHONY: all test clean + +all: $(TESTS) + +test: $(TESTS) + @failed=0; \ + for t in $(TESTS); do \ + echo "--- Running $$t ---"; \ + ./$$t || failed=1; \ + echo; \ + done; \ + if [ $$failed -eq 0 ]; then echo "All tests passed."; else echo "SOME TESTS FAILED"; exit 1; fi + +test_seqlock: test_seqlock.cpp $(PARENT_OBJS) + $(CXX) $(CXXFLAGS) -o $@ $^ + +test_measurement: test_measurement.cpp $(PARENT_OBJS) + $(CXX) $(CXXFLAGS) -o $@ $^ + +test_regression: test_regression.cpp $(PARENT_OBJS) + $(CXX) $(CXXFLAGS) -o $@ $^ + +$(PARENT_OBJS): + $(MAKE) -C .. $(notdir $@) + +clean: + rm -f $(TESTS) diff --git a/rpd_tracer/tests/test_measurement.cpp b/rpd_tracer/tests/test_measurement.cpp new file mode 100644 index 0000000..ab2b5e6 --- /dev/null +++ b/rpd_tracer/tests/test_measurement.cpp @@ -0,0 +1,91 @@ +/************************************************************************** + * Copyright (c) 2024 Advanced Micro Devices, Inc. + **************************************************************************/ +#include "Firefly.h" + +#include +#include +#include +#include + +using namespace rpdtracer::firefly; + +static Measurement make_measurement(int64_t offset, char node) +{ + Measurement m{}; + m.offset = offset; + m.node[0] = node; + m.node[1] = '\0'; + return m; +} + +static void test_basic_insert() +{ + MeasurementBuffer buf(100); + assert(buf.count == 0); + + { + std::lock_guard lock(buf.mutex); + buf.measurements[buf.count] = make_measurement(42, 'A'); + ++buf.count; + } + + assert(buf.count == 1); + assert(buf.measurements[0].offset == 42); + assert(buf.measurements[0].node[0] == 'A'); + fprintf(stderr, " PASS: basic insert\n"); +} + +static void test_concurrent_inserts() +{ + const int threads = 8; + const int per_thread = 1000; + MeasurementBuffer buf(threads * per_thread); + + auto worker = [&](int id) { + for (int i = 0; i < per_thread; ++i) { + Measurement m = make_measurement(id * per_thread + i, 'A'); + std::lock_guard lock(buf.mutex); + if (buf.count < static_cast(buf.measurements.size())) { + buf.measurements[buf.count] = m; + ++buf.count; + } + } + }; + + std::vector workers; + for (int i = 0; i < threads; ++i) + workers.emplace_back(worker, i); + for (auto& w : workers) + w.join(); + + assert(buf.count == threads * per_thread); + fprintf(stderr, " PASS: concurrent inserts (%d threads x %d = %d total)\n", + threads, per_thread, buf.count); +} + +static void test_capacity_limit() +{ + MeasurementBuffer buf(10); + + for (int i = 0; i < 20; ++i) { + std::lock_guard lock(buf.mutex); + if (buf.count < static_cast(buf.measurements.size())) { + buf.measurements[buf.count] = make_measurement(i, 'B'); + ++buf.count; + } + } + + assert(buf.count == 10); + fprintf(stderr, " PASS: capacity limit (10 stored, 10 rejected)\n"); +} + +int main() +{ + fprintf(stderr, "test_measurement:\n"); + test_basic_insert(); + test_concurrent_inserts(); + test_capacity_limit(); + fprintf(stderr, " All measurement tests passed.\n"); + return 0; +} diff --git a/rpd_tracer/tests/test_regression.cpp b/rpd_tracer/tests/test_regression.cpp new file mode 100644 index 0000000..7ae020f --- /dev/null +++ b/rpd_tracer/tests/test_regression.cpp @@ -0,0 +1,153 @@ +/************************************************************************** + * Copyright (c) 2024 Advanced Micro Devices, Inc. + **************************************************************************/ +#include "Firefly.h" + +#include +#include +#include + +using namespace rpdtracer::firefly; + +static Measurement make_timed(timestamp_t ts, int64_t offset, char node) +{ + Measurement m{}; + m.timestampNs = ts; + m.offset = offset; + m.node[0] = node; + m.node[1] = '\0'; + return m; +} + +static void test_constant_offset() +{ + const int N = 100; + Measurement data[N]; + for (int i = 0; i < N; ++i) + data[i] = make_timed(1000000000ULL + i * 1000000ULL, 5000, 'A'); + + MeasurementAnalysis a = read_latest_measurements("A", 1000, data, N); + assert(a.samples.size() == static_cast(N)); + assert(a.averageOffset == 5000); + assert(std::fabs(a.driftRate) < 1e-12); + fprintf(stderr, " PASS: constant offset (avg=%lld, drift=%.2e)\n", + static_cast(a.averageOffset), a.driftRate); +} + +static void test_known_drift() +{ + const int N = 100; + const double drift_per_ns = 0.0001; + Measurement data[N]; + + timestamp_t t0 = 1000000000ULL; + for (int i = 0; i < N; ++i) { + timestamp_t ts = t0 + i * 1000000ULL; + int64_t offset = static_cast(1000 + drift_per_ns * (ts - t0)); + data[i] = make_timed(ts, offset, 'A'); + } + + MeasurementAnalysis a = read_latest_measurements("A", 1000, data, N); + assert(a.samples.size() == static_cast(N)); + double drift_error = std::fabs(a.driftRate - drift_per_ns); + assert(drift_error < 1e-8); + fprintf(stderr, " PASS: known drift (expected=%.4e, got=%.4e, err=%.2e)\n", + drift_per_ns, a.driftRate, drift_error); +} + +static void test_role_filter() +{ + Measurement data[4]; + data[0] = make_timed(1000000000ULL, 100, 'A'); + data[1] = make_timed(1000000001ULL, 200, 'B'); + data[2] = make_timed(1000000002ULL, 300, 'A'); + data[3] = make_timed(1000000003ULL, 400, 'B'); + + MeasurementAnalysis a = read_latest_measurements("A", 1000, data, 4); + assert(a.samples.size() == 2); + assert(a.samples[0].offset == 100); + assert(a.samples[1].offset == 300); + + MeasurementAnalysis b = read_latest_measurements("B", 1000, data, 4); + assert(b.samples.size() == 2); + assert(b.samples[0].offset == 200); + assert(b.samples[1].offset == 400); + + fprintf(stderr, " PASS: role filter\n"); +} + +static void test_empty_input() +{ + MeasurementAnalysis a = read_latest_measurements("A", 1000, nullptr, 0); + assert(a.samples.empty()); + assert(a.averageOffset == 0); + assert(a.driftRate == 0.0); + fprintf(stderr, " PASS: empty input\n"); +} + +static void test_single_measurement() +{ + Measurement data[1]; + data[0] = make_timed(1000000000ULL, 42, 'A'); + + MeasurementAnalysis a = read_latest_measurements("A", 1000, data, 1); + assert(a.samples.size() == 1); + assert(a.averageOffset == 42); + assert(a.driftRate == 0.0); + fprintf(stderr, " PASS: single measurement\n"); +} + +static void test_window_truncation() +{ + const int N = 100; + Measurement data[N]; + for (int i = 0; i < N; ++i) + data[i] = make_timed(1000000000ULL + i * 1000000ULL, i, 'A'); + + MeasurementAnalysis a = read_latest_measurements("A", 10, data, N); + assert(a.samples.size() == 10); + assert(a.samples[0].offset == 90); + assert(a.samples[9].offset == 99); + fprintf(stderr, " PASS: window truncation (100 -> 10)\n"); +} + +static void test_drift_clamp() +{ + // Build a buffer with measurements that produce drift > 500 ppm + const int N = 100; + const double bad_drift = 0.01; // 10,000,000 ppm — way over 500 ppm + MeasurementBuffer buf(N); + + timestamp_t t0 = 1000000000ULL; + for (int i = 0; i < N; ++i) { + timestamp_t ts = t0 + i * 1000000ULL; + int64_t offset = static_cast(bad_drift * (ts - t0)); + buf.measurements[i] = make_timed(ts, offset, 'A'); + } + buf.count = N; + + // Reset g_svcState + g_pSvcState->sequence.store(0, std::memory_order_relaxed); + g_pSvcState->offset = 0; + g_pSvcState->drift = 0.0; + + firefly_run("A", buf); + + // Drift should have been clamped to 0 + assert(std::fabs(g_pSvcState->drift) < 1e-15); + fprintf(stderr, " PASS: drift clamp (bad drift=%.2e clamped to 0)\n", bad_drift); +} + +int main() +{ + fprintf(stderr, "test_regression:\n"); + test_constant_offset(); + test_known_drift(); + test_role_filter(); + test_empty_input(); + test_single_measurement(); + test_window_truncation(); + test_drift_clamp(); + fprintf(stderr, " All regression tests passed.\n"); + return 0; +} diff --git a/rpd_tracer/tests/test_seqlock.cpp b/rpd_tracer/tests/test_seqlock.cpp new file mode 100644 index 0000000..c5f7926 --- /dev/null +++ b/rpd_tracer/tests/test_seqlock.cpp @@ -0,0 +1,103 @@ +/************************************************************************** + * Copyright (c) 2024 Advanced Micro Devices, Inc. + **************************************************************************/ +#include "Utility.h" +#include "Firefly.h" + +#include +#include +#include +#include +#include +#include + +using namespace rpdtracer; +using namespace rpdtracer::firefly; + +static void test_seq_zero_returns_zero() +{ + g_pSvcState->sequence.store(0, std::memory_order_relaxed); + g_pSvcState->offset = 0; + g_pSvcState->drift = 0.0; + + int64_t result = svc_read_offset_ns(1000000000ULL); + assert(result == 0); + fprintf(stderr, " PASS: seq==0 fast path returns 0\n"); +} + +static void test_read_back_written_values() +{ + svc_update_ns(g_pSvcState, 5000, 0.0); + + int64_t result = svc_read_offset_ns(1000000000ULL); + assert(result == 5000); + fprintf(stderr, " PASS: offset-only read back\n"); +} + +static void test_drift_correction() +{ + svc_update_ns(g_pSvcState, 1000, 0.001); + + // Read referenceTime that svc_update_ns just wrote + timestamp_t ref_ns = timespec_to_ns(g_pSvcState->referenceTime); + timestamp_t now_ns = ref_ns + 1000000; + int64_t result = svc_read_offset_ns(now_ns); + + // offset + drift * elapsed = 1000 + 0.001 * 1000000 = 2000 + assert(result > 1900 && result < 2100); + fprintf(stderr, " PASS: drift correction (result=%lld)\n", static_cast(result)); +} + +static void test_concurrent_no_torn_reads() +{ + std::atomic done{false}; + std::atomic errors{0}; + + // Writer alternates between two distinct states with pauses + // so readers can complete seqlock reads between writes + std::thread writer([&]() { + for (int i = 0; i < 10000 && !done.load(std::memory_order_relaxed); ++i) { + if (i % 2 == 0) + svc_update_ns(g_pSvcState, 0, 0.0); + else + svc_update_ns(g_pSvcState, 1000000, 0.0); + std::this_thread::sleep_for(std::chrono::microseconds(10)); + } + done.store(true, std::memory_order_relaxed); + }); + + // Readers verify result is one of the two expected values + auto reader_fn = [&]() { + while (!done.load(std::memory_order_relaxed)) { + timespec ts; + clock_gettime(CLOCK_MONOTONIC, &ts); + timestamp_t now = timespec_to_ns(ts); + int64_t result = svc_read_offset_ns(now); + // With drift=0, result should be exactly 0 or 1000000 + if (result != 0 && result != 1000000) + errors.fetch_add(1, std::memory_order_relaxed); + } + }; + + std::vector readers; + for (int i = 0; i < 4; ++i) + readers.emplace_back(reader_fn); + + writer.join(); + for (auto& r : readers) + r.join(); + + assert(errors.load() == 0); + fprintf(stderr, " PASS: concurrent seqlock stress (4 readers, no torn reads)\n"); +} + +int main() +{ + fprintf(stderr, "test_seqlock:\n"); + test_seq_zero_returns_zero(); + test_read_back_written_values(); + test_drift_correction(); + test_concurrent_no_torn_reads(); + fprintf(stderr, " All seqlock tests passed.\n"); + return 0; +} diff --git a/tools/compare_rpds.py b/tools/compare_rpds.py new file mode 100644 index 0000000..ca0d5ca --- /dev/null +++ b/tools/compare_rpds.py @@ -0,0 +1,143 @@ +import os +import sqlite3 +import pandas as pd +import sys + +def extract_rccl_kernels_from_file(db_path): + """ + Extract RCCL kernels from a single SQLite file. + """ + try: + # Connect to the SQLite database + conn = sqlite3.connect(db_path) + cursor = conn.cursor() + + # Query to extract RCCL kernels + query = """ + SELECT id, gpuId, queueId, sequenceId, start, end, duration, stream, + gridX, gridY, gridZ, workgroupX, workgroupY, workgroupZ, + groupSegmentSize, privateSegmentSize, kernelName + FROM kernel + WHERE kernelName LIKE "%rccl%" + """ + cursor.execute(query) + + # Fetch all results + rows = cursor.fetchall() + + # Define column names + columns = [desc[0] for desc in cursor.description] + + # Create a pandas DataFrame + df = pd.DataFrame(rows, columns=columns) + + # Close the connection + conn.close() + + return df + except Exception as e: + print(f"Error processing file {db_path}: {e}") + return pd.DataFrame() # Return an empty DataFrame on error + +def process_rpd_files(input_dir): + """ + Process all .rpd files in the input directory and return a dictionary of DataFrames. + """ + data_frames = {} + + # Iterate over all files in the input directory + for root, _, files in os.walk(input_dir): + for file in files: + if file.endswith(".rpd"): + file_path = os.path.join(root, file) + print(f"Processing file: {file_path}") + + # Extract RCCL kernels from the file + df = extract_rccl_kernels_from_file(file_path) + if not df.empty: + # Store the DataFrame in the dictionary with the file path as the key + data_frames[file_path] = df + + return data_frames + +def compare_data_frames(data_frames): + """ + Compare multiple DataFrames to check for consistency in RCCL calls. + """ + if not data_frames: + print("No data frames to compare.") + return + + # Extract all DataFrames and their file paths + file_paths = list(data_frames.keys()) + dfs = list(data_frames.values()) + + # Check if all DataFrames have the same number of RCCL calls + num_calls = [len(df) for df in dfs] + if len(set(num_calls)) != 1: + print("Mismatch in the number of RCCL calls across files:") + for path, count in zip(file_paths, num_calls): + print(f"{path}: {count} calls") + return + + print("All files have the same number of RCCL calls.") + + # Check for overlapping start and end times for corresponding RCCL calls + error_count = 0 + total_overlap_percentage = 0 + num_calls = len(dfs[0]) # Number of RCCL calls in each DataFrame + for i in range(num_calls): + for j in range(len(dfs)): + for k in range(j + 1, len(dfs)): + start_j = dfs[j].iloc[i]["start"] + end_j = dfs[j].iloc[i]["end"] + start_k = dfs[k].iloc[i]["start"] + end_k = dfs[k].iloc[i]["end"] + + # Calculate overlap duration + overlap_start = max(start_j, start_k) + overlap_end = min(end_j, end_k) + overlap_duration = max(0, overlap_end - overlap_start) + + # Calculate the total duration of the two calls + total_duration = max(end_j, end_k) - min(start_j, start_k) + + # Calculate overlap percentage + if total_duration > 0: + overlap_percentage = (overlap_duration / total_duration) * 100 + total_overlap_percentage += overlap_percentage + + # Check if there is no overlap + if overlap_duration == 0: + error_count += 1 + print(f"Error in RCCL call {i} between files {file_paths[j]} and {file_paths[k]}: No overlap in start/end times.") + + # Calculate overall average overlap percentage + if num_calls > 0 and len(dfs) > 1: + total_comparisons = num_calls * (len(dfs) * (len(dfs) - 1)) / 2 + average_overlap_percentage = total_overlap_percentage / total_comparisons + print(f"Overall average overlap percentage: {average_overlap_percentage:.2f}%") + + print(f"Total number of errors: {error_count}") + +if __name__ == "__main__": + # Input directory containing .rpd files + # input_directory = "/home/docker/jax-llm-examples/llama3/clock_sync_results/with_sync" + # input_directory = "/home/docker/jax-llm-examples/llama3/clock_sync_results/without_sync" + + if len(sys.argv) != 2: + print("Usage: python compare_rpds.py ") + sys.exit(1) + + input_directory = sys.argv[1] + # Process the .rpd files and extract RCCL kernels + data_frames = process_rpd_files(input_directory) + + compare_data_frames(data_frames) + # Save each DataFrame to a separate CSV file + for file_path, df in data_frames.items(): + output_csv = f"{file_path}.rccl_kernels.csv" + df.to_csv(output_csv, index=False) + print(f"RCCL kernel data for {file_path} saved to {output_csv}") + +# python3 /home/docker/rocmProfileData/tools/compare_rpds.py \ No newline at end of file diff --git a/tools/multiplerpd2tracing.py b/tools/multiplerpd2tracing.py new file mode 100644 index 0000000..efddb66 --- /dev/null +++ b/tools/multiplerpd2tracing.py @@ -0,0 +1,353 @@ +#!/usr/bin/env python3 + +################################################################################ +# Copyright (c) 2021 - 2023 Advanced Micro Devices, Inc. All rights reserved. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +# THE SOFTWARE. +################################################################################ + +# +# Format sqlite trace data as json for chrome:tracing - Multi-file merger +# + +import sys +import os +import csv +import re +import sqlite3 +from collections import defaultdict +from datetime import datetime +import argparse + +parser = argparse.ArgumentParser(description='convert and merge multiple RPD files to json for chrome tracing') +parser.add_argument('input_rpd', type=str, nargs='+', help="input rpd db files (can specify multiple)") +parser.add_argument('output_json', type=str, nargs='?', help="chrome tracing json output") +parser.add_argument('--start', type=str, help="start time - default us or percentage %%. Number only is interpreted as us. Number with %% is interpreted as percentage") +parser.add_argument('--end', type=str, help="end time - default us or percentage %%. See help for --start") +parser.add_argument('--format', type=str, default="object", help="chrome trace format, array or object") +args = parser.parse_args() + +if args.output_json is None: + import pathlib + if len(args.input_rpd) == 1: + args.output_json = pathlib.PurePath(args.input_rpd[0]).with_suffix(".json") + else: + args.output_json = "merged_trace.json" + +outfile = open(args.output_json, 'w', encoding="utf-8") + +if args.format == "object": + outfile.write("{\"traceEvents\": ") + +outfile.write("[ {}\n") + +def process_rpd_file(rpd_file, node_id, outfile, args): + """Process a single RPD file with node prefix""" + + print(f"\n{'='*80}") + print(f"Processing Node {node_id}: {rpd_file}") + print(f"{'='*80}") + + connection = sqlite3.connect(rpd_file) + + # Node prefix for display names + node_prefix = f"Node{node_id}" + + # PID offset to avoid conflicts between files + pid_offset = node_id * 100000 + + # GPU metadata + for row in connection.execute("select distinct gpuId from rocpd_op"): + try: + gpu_pid = row[0] + pid_offset + outfile.write(",{\"name\": \"process_name\", \"ph\": \"M\", \"pid\":\"%s\",\"args\":{\"name\":\"%s GPU%s\"}}\n"%(gpu_pid, node_prefix, row[0])) + outfile.write(",{\"name\": \"process_sort_index\", \"ph\": \"M\", \"pid\":\"%s\",\"args\":{\"sort_index\":\"%s\"}}\n"%(gpu_pid, gpu_pid + 1000000)) + except ValueError: + pass + + # Thread metadata for HIP APIs + for row in connection.execute("select distinct pid, tid from rocpd_api"): + try: + adj_pid = row[0] + pid_offset + outfile.write(',{"name":"thread_name","ph":"M","pid":"%s","tid":"%s","args":{"name":"%s Hip %s"}}\n'%(adj_pid, row[1], node_prefix, row[1])) + outfile.write(',{"name":"thread_sort_index","ph":"M","pid":"%s","tid":"%s","args":{"sort_index":"%s"}}\n'%(adj_pid, row[1], row[1] * 2)) + except ValueError: + pass + + # Thread metadata for HSA APIs + try: + for row in connection.execute("select distinct pid, tid from rocpd_hsaApi"): + try: + adj_pid = row[0] + pid_offset + outfile.write(',{"name":"thread_name","ph":"M","pid":"%s","tid":"%s","args":{"name":"%s HSA %s"}}\n'%(adj_pid, row[1], node_prefix, row[1])) + outfile.write(',{"name":"thread_sort_index","ph":"M","pid":"%s","tid":"%s","args":{"sort_index":"%s"}}\n'%(adj_pid, row[1], row[1] * 2 - 1)) + except ValueError: + pass + except: + pass + + # Time range calculation + rangeStringApi = "" + rangeStringOp = "" + rangeStringMonitor = "" + min_time = connection.execute("select MIN(start) from rocpd_api;").fetchall()[0][0] + max_time = connection.execute("select MAX(end) from rocpd_api;").fetchall()[0][0] + + if min_time is None: + print(f"Warning: Trace file {rpd_file} is empty, skipping...") + connection.close() + return + + print(f"Timestamps for {node_prefix}:") + print(f"\t first: \t{min_time/1000} us") + print(f"\t last: \t{max_time/1000} us") + print(f"\t duration: \t{(max_time-min_time) / 1000000000} seconds") + + start_time = min_time/1000 + end_time = max_time/1000 + + if args.start: + if "%" in args.start: + start_time = ( (max_time - min_time) * ( int( args.start.replace("%","") )/100 ) + min_time )/1000 + else: + start_time = int(args.start) + rangeStringApi = "where rocpd_api.start/1000 >= %s"%(start_time) + rangeStringOp = "where rocpd_op.start/1000 >= %s"%(start_time) + rangeStringMonitor = "where start/1000 >= %s"%(start_time) + + if args.end: + if "%" in args.end: + end_time = ( (max_time - min_time) * ( int( args.end.replace("%","") )/100 ) + min_time )/1000 + else: + end_time = int(args.end) + + rangeStringApi = rangeStringApi + " and rocpd_api.start/1000 <= %s"%(end_time) if args.start != None else "where rocpd_api.start/1000 <= %s"%(end_time) + rangeStringOp = rangeStringOp + " and rocpd_op.start/1000 <= %s"%(end_time) if args.start != None else "where rocpd_op.start/1000 <= %s"%(end_time) + rangeStringOp = rangeStringOp + " and rocpd_op.start/1000 <= %s"%(end_time) if args.start != None else "where rocpd_op.start/1000 <= %s"%(end_time) + rangeStringMonitor = rangeStringMonitor + " and start/1000 <= %s"%(end_time) if args.start != None else "where start/1000 <= %s"%(end_time) + + print(f"\nFilter for {node_prefix}: {rangeStringApi}") + print(f"Output duration: {(end_time-start_time)/1000000} seconds") + + # Output Ops + for row in connection.execute("select A.string as optype, B.string as description, gpuId, queueId, rocpd_op.start/1000.0, (rocpd_op.end-rocpd_op.start) / 1000.0 from rocpd_op INNER JOIN rocpd_string A on A.id = rocpd_op.opType_id INNER Join rocpd_string B on B.id = rocpd_op.description_id %s"%(rangeStringOp)): + try: + name = row[0] if len(row[1])==0 else row[1] + gpu_pid = row[2] + pid_offset + outfile.write(",{\"pid\":\"%s\",\"tid\":\"%s\",\"name\":\"%s\",\"ts\":\"%s\",\"dur\":\"%s\",\"ph\":\"X\",\"args\":{\"desc\":\"%s\"}}\n"%(gpu_pid, row[3], name, row[4], row[5], row[0])) + except ValueError: + pass + + # Output Graph executions on GPU + try: + for row in connection.execute('select graphExec, gpuId, queueId, min(start)/1000.0, (max(end)-min(start))/1000.0, count(*) from rocpd_graphLaunchapi A join rocpd_api_ops B on B.api_id = A.api_ptr_id join rocpd_op C on C.id = B.op_id %s group by api_ptr_id'%(rangeStringMonitor)): + try: + gpu_pid = row[1] + pid_offset + outfile.write(",{\"pid\":\"%s\",\"tid\":\"%s\",\"name\":\"%s\",\"ts\":\"%s\",\"dur\":\"%s\",\"ph\":\"X\",\"args\":{\"kernels\":\"%s\"}}\n"%(gpu_pid, row[2], f'Graph {row[0]}', row[3], row[4], row[5])) + except ValueError: + pass + except: + pass + + # Output APIs + for row in connection.execute("select A.string as apiName, B.string as args, pid, tid, rocpd_api.start/1000.0, (rocpd_api.end-rocpd_api.start) / 1000.0, (rocpd_api.end != rocpd_api.start) as has_duration from rocpd_api INNER JOIN rocpd_string A on A.id = rocpd_api.apiName_id INNER Join rocpd_string B on B.id = rocpd_api.args_id %s order by rocpd_api.id"%(rangeStringApi)): + try: + adj_pid = row[2] + pid_offset + if row[0]=="UserMarker": + if row[6] == 0: # instantaneous "mark" messages + outfile.write(",{\"pid\":\"%s\",\"tid\":\"%s\",\"name\":\"%s\",\"ts\":\"%s\",\"ph\":\"i\",\"s\":\"p\",\"args\":{\"desc\":\"%s\"}}\n"%(adj_pid, row[3], row[1].replace('"',''), row[4], row[1].replace('"',''))) + else: + outfile.write(",{\"pid\":\"%s\",\"tid\":\"%s\",\"name\":\"%s\",\"ts\":\"%s\",\"dur\":\"%s\",\"ph\":\"X\",\"args\":{\"desc\":\"%s\"}}\n"%(adj_pid, row[3], row[1].replace('"',''), row[4], row[5], row[1].replace('"',''))) + else: + outfile.write(",{\"pid\":\"%s\",\"tid\":\"%s\",\"name\":\"%s\",\"ts\":\"%s\",\"dur\":\"%s\",\"ph\":\"X\",\"args\":{\"desc\":\"%s\"}}\n"%(adj_pid, row[3], row[0], row[4], row[5], row[1].replace('"','').replace('\t',''))) + except ValueError: + pass + + # Output api->op linkage + for row in connection.execute("select rocpd_api_ops.id, pid, tid, gpuId, queueId, rocpd_api.end/1000.0 - 2, rocpd_op.start/1000.0 from rocpd_api_ops INNER JOIN rocpd_api on rocpd_api_ops.api_id = rocpd_api.id INNER JOIN rocpd_op on rocpd_api_ops.op_id = rocpd_op.id %s"%(rangeStringApi)): + try: + fromtime = row[5] if row[5] < row[6] else row[6] + adj_pid = row[1] + pid_offset + gpu_pid = row[3] + pid_offset + # Use unique IDs per node to avoid conflicts + link_id = row[0] + (node_id * 10000000) + outfile.write(",{\"pid\":\"%s\",\"tid\":\"%s\",\"cat\":\"api_op\",\"name\":\"api_op\",\"ts\":\"%s\",\"id\":\"%s\",\"ph\":\"s\"}\n"%(adj_pid, row[2], fromtime, link_id)) + outfile.write(",{\"pid\":\"%s\",\"tid\":\"%s\",\"cat\":\"api_op\",\"name\":\"api_op\",\"ts\":\"%s\",\"id\":\"%s\",\"ph\":\"f\", \"bp\":\"e\"}\n"%(gpu_pid, row[4], row[6], link_id)) + except ValueError: + pass + + # Output HSA APIs + try: + for row in connection.execute("select A.string as apiName, B.string as args, pid, tid, rocpd_hsaApi.start/1000.0, (rocpd_hsaApi.end-rocpd_hsaApi.start) / 1000.0 from rocpd_hsaApi INNER JOIN rocpd_string A on A.id = rocpd_hsaApi.apiName_id INNER Join rocpd_string B on B.id = rocpd_hsaApi.args_id %s order by rocpd_hsaApi.id"%(rangeStringApi)): + try: + adj_pid = row[2] + pid_offset + outfile.write(",{\"pid\":\"%s\",\"tid\":\"%s\",\"name\":\"%s\",\"ts\":\"%s\",\"dur\":\"%s\",\"ph\":\"X\",\"args\":{\"desc\":\"%s\"}}\n"%(adj_pid, row[3]+1, row[0], row[4], row[5], row[1].replace('"',''))) + except ValueError: + pass + except: + pass + + # + # Counters + # + + # Counters should extend to the last event in the trace + # + # Counters + # + + # Counters should extend to the last event in the trace + T_end = 0 + for row in connection.execute("SELECT max(end)/1000 from (SELECT end from rocpd_api UNION ALL SELECT end from rocpd_op)"): + T_end = int(row[0]) + if args.end: + T_end = end_time + + # Loop over GPU for per-gpu counters + gpuIdsPresent = [] + for row in connection.execute("SELECT DISTINCT gpuId FROM rocpd_op"): + gpuIdsPresent.append(row[0]) + + for gpuId in gpuIdsPresent: + gpu_pid = gpuId + pid_offset + + # Create the queue depth counter + depth = 0 + idle = 1 + for row in connection.execute("select * from (select rocpd_api.start/1000.0 as ts, \"1\" from rocpd_api_ops INNER JOIN rocpd_api on rocpd_api_ops.api_id = rocpd_api.id INNER JOIN rocpd_op on rocpd_api_ops.op_id = rocpd_op.id AND rocpd_op.gpuId = %s %s UNION ALL select rocpd_op.end/1000.0, \"-1\" from rocpd_api_ops INNER JOIN rocpd_api on rocpd_api_ops.api_id = rocpd_api.id INNER JOIN rocpd_op on rocpd_api_ops.op_id = rocpd_op.id AND rocpd_op.gpuId = %s %s) order by ts"%(gpuId, rangeStringOp, gpuId, rangeStringOp)): + try: + if idle and int(row[1]) > 0: + idle = 0 + outfile.write(',{"pid":"%s","name":"Idle","ph":"C","ts":%s,"args":{"idle":%s}}\n'%(gpu_pid, row[0], idle)) + if depth == 1 and int(row[1]) < 0: + idle = 1 + outfile.write(',{"pid":"%s","name":"Idle","ph":"C","ts":%s,"args":{"idle":%s}}\n'%(gpu_pid, row[0], idle)) + depth = depth + int(row[1]) + outfile.write(',{"pid":"%s","name":"QueueDepth","ph":"C","ts":%s,"args":{"depth":%s}}\n'%(gpu_pid, row[0], depth)) + except ValueError: + pass + if T_end > 0: + outfile.write(',{"pid":"%s","name":"Idle","ph":"C","ts":%s,"args":{"idle":%s}}\n'%(gpu_pid, T_end, idle)) + outfile.write(',{"pid":"%s","name":"QueueDepth","ph":"C","ts":%s,"args":{"depth":%s}}\n'%(gpu_pid, T_end, depth)) + + # Create SMI counters + try: + for row in connection.execute("select deviceId, monitorType, start/1000.0, value from rocpd_monitor %s"%(rangeStringMonitor)): + device_pid = row[0] + pid_offset + outfile.write(',{"pid":"%s","name":"%s","ph":"C","ts":%s,"args":{"%s":%s}}\n'%(device_pid, row[1], row[2], row[1], row[3])) + # Output the endpoints of the last range + for row in connection.execute("select distinct deviceId, monitorType, max(end)/1000.0, value from rocpd_monitor %s group by deviceId, monitorType"%(rangeStringMonitor)): + device_pid = row[0] + pid_offset + outfile.write(',{"pid":"%s","name":"%s","ph":"C","ts":%s,"args":{"%s":%s}}\n'%(device_pid, row[1], row[2], row[1], row[3])) + except: + print(f"Did not find SMI data for {node_prefix}") + + # Create "faux calling stack frame" on gpu ops traces + stacks = {} # Call stacks built from UserMarker entries. Key is 'pid,tid' + currentFrame = {} # "Current GPU frame" (id, name, start, end). Key is 'pid,tid' + + class GpuFrame: + def __init__(self): + self.id = 0 + self.name = '' + self.start = 0 + self.end = 0 + self.gpus = [] + self.totalOps = 0 + + for row in connection.execute("SELECT '0', start/1000.0, pid, tid, B.string as label, '','','', '' from rocpd_api INNER JOIN rocpd_string A on A.id = rocpd_api.apiName_id AND A.string = 'UserMarker' INNER JOIN rocpd_string B on B.id = rocpd_api.args_id AND rocpd_api.start/1000.0 != rocpd_api.end/1000.0 %s UNION ALL SELECT '1', end/1000.0, pid, tid, B.string as label, '','','', '' from rocpd_api INNER JOIN rocpd_string A on A.id = rocpd_api.apiName_id AND A.string = 'UserMarker' INNER JOIN rocpd_string B on B.id = rocpd_api.args_id AND rocpd_api.start/1000.0 != rocpd_api.end/1000.0 %s UNION ALL SELECT '2', rocpd_api.start/1000.0, pid, tid, '' as label, gpuId, queueId, rocpd_op.start/1000.0, rocpd_op.end/1000.0 from rocpd_api_ops INNER JOIN rocpd_api ON rocpd_api_ops.api_id = rocpd_api.id INNER JOIN rocpd_op ON rocpd_api_ops.op_id = rocpd_op.id %s ORDER BY start/1000.0 asc"%(rangeStringApi, rangeStringApi, rangeStringApi)): + try: + key = (row[2], row[3]) # Key is 'pid,tid' + if row[0] == '0': # Frame start + if key not in stacks: + stacks[key] = [] + stack = stacks[key].append((row[1], row[4])) + + elif row[0] == '1': # Frame end + if key in stacks and len(stacks[key]) > 0: + completed = stacks[key].pop() + + elif row[0] == '2': # API + Op + if key in stacks and len(stacks[key]) > 0: + frame = stacks[key][-1] + gpuFrame = None + if key not in currentFrame: # First op under the current api frame + gpuFrame = GpuFrame() + gpuFrame.id = frame[0] + gpuFrame.name = frame[1] + gpuFrame.start = row[7] + gpuFrame.end = row[8] + gpuFrame.gpus.append((row[5] + pid_offset, row[6])) + gpuFrame.totalOps = 1 + else: + gpuFrame = currentFrame[key] + # Another op under the same frame -> union them (but only if they are close together) + if gpuFrame.id == frame[0] and gpuFrame.name == frame[1] and (abs(row[7] - gpuFrame.end) < 200 or abs(gpuFrame.start - row[8]) < 200): + if row[7] < gpuFrame.start: gpuFrame.start = row[7] + if row[8] > gpuFrame.end: gpuFrame.end = row[8] + if (row[5] + pid_offset, row[6]) not in gpuFrame.gpus: + gpuFrame.gpus.append((row[5] + pid_offset, row[6])) + gpuFrame.totalOps = gpuFrame.totalOps + 1 + + else: # This is a new frame - dump the last and make new + gpuFrame = currentFrame[key] + for dest in gpuFrame.gpus: + outfile.write(',{"pid":"%s","tid":"%s","name":"%s","ts":"%s","dur":"%s","ph":"X","args":{"desc":"%s"}}\n'%(dest[0], dest[1], gpuFrame.name.replace('"',''), gpuFrame.start - 1, gpuFrame.end - gpuFrame.start + 1, f"UserMarker frame: {gpuFrame.totalOps} ops")) + currentFrame.pop(key) + + # make the first op under the new frame + gpuFrame = GpuFrame() + gpuFrame.id = frame[0] + gpuFrame.name = frame[1] + gpuFrame.start = row[7] + gpuFrame.end = row[8] + gpuFrame.gpus.append((row[5] + pid_offset, row[6])) + gpuFrame.totalOps = 1 + + currentFrame[key] = gpuFrame + + except ValueError: + pass + + connection.close() + print(f"Finished processing {node_prefix}") + + +# Main execution: Process all input files +for idx, rpd_file in enumerate(args.input_rpd): + node_id = idx # Node 0, Node 1, etc. + try: + process_rpd_file(rpd_file, node_id, outfile, args) + except Exception as e: + print(f"Error processing {rpd_file}: {e}") + import traceback + traceback.print_exc() + +# Close the JSON output +outfile.write("]\n") + +if args.format == "object": + outfile.write("} \n") + +outfile.close() + +print(f"\n{'='*80}") +print(f"Merged trace written to: {args.output_json}") +print(f"Total files processed: {len(args.input_rpd)}") +print(f"{'='*80}")