fixed npkit size to never be a negative number (#779)
[ROCm/rccl commit: 9bdf6797a5]
This commit is contained in:
@@ -30,20 +30,20 @@ class NpKit {
|
||||
|
||||
static NpKitEventCollectContext* GetGpuEventCollectContexts();
|
||||
|
||||
static inline __device__ void CollectGpuEvent(uint8_t type, uint32_t size, uint32_t rsvd, uint64_t timestamp,
|
||||
static inline __device__ void CollectGpuEvent(uint8_t type, int64_t size, uint32_t rsvd, uint64_t timestamp,
|
||||
NpKitEventCollectContext* ctx) {
|
||||
uint64_t event_buffer_head = ctx->event_buffer_head;
|
||||
if (event_buffer_head < kMaxNumGpuEventsPerBuffer) {
|
||||
NpKitEvent& event = ctx->event_buffer[event_buffer_head];
|
||||
event.fields.type = type;
|
||||
event.fields.size = size;
|
||||
event.fields.size = size < 0 ? 0 : size;
|
||||
event.fields.rsvd = rsvd;
|
||||
event.fields.timestamp = timestamp;
|
||||
ctx->event_buffer_head++;
|
||||
}
|
||||
}
|
||||
|
||||
static void CollectCpuEvent(uint8_t type, uint32_t size, uint32_t rsvd, uint64_t timestamp, int channel_id);
|
||||
static void CollectCpuEvent(uint8_t type, int64_t size, uint32_t rsvd, uint64_t timestamp, int channel_id);
|
||||
|
||||
static uint64_t *GetCpuTimestamp();
|
||||
|
||||
|
||||
@@ -14,7 +14,7 @@ union NpKitEvent {
|
||||
uint64_t bits[2];
|
||||
struct {
|
||||
uint64_t type : 8;
|
||||
uint64_t size : 32;
|
||||
uint32_t size : 32;
|
||||
uint64_t rsvd : 24;
|
||||
uint64_t timestamp;
|
||||
} fields;
|
||||
|
||||
@@ -160,12 +160,12 @@ NpKitEventCollectContext* NpKit::GetGpuEventCollectContexts() {
|
||||
return gpu_collect_contexts_;
|
||||
}
|
||||
|
||||
void NpKit::CollectCpuEvent(uint8_t type, uint32_t size, uint32_t rsvd, uint64_t timestamp, int channel_id) {
|
||||
void NpKit::CollectCpuEvent(uint8_t type, int64_t size, uint32_t rsvd, uint64_t timestamp, int channel_id) {
|
||||
uint64_t event_buffer_head = cpu_collect_contexts_[channel_id].event_buffer_head;
|
||||
if (event_buffer_head < kMaxNumCpuEventsPerBuffer) {
|
||||
NpKitEvent& event = cpu_collect_contexts_[channel_id].event_buffer[event_buffer_head];
|
||||
event.fields.type = type;
|
||||
event.fields.size = size;
|
||||
event.fields.size = size < 0 ? 0 : size;
|
||||
event.fields.rsvd = rsvd;
|
||||
event.fields.timestamp = timestamp;
|
||||
cpu_collect_contexts_[channel_id].event_buffer_head++;
|
||||
|
||||
Reference in New Issue
Block a user