codekingpro/portable-devtools
114k
1# -------------------------------------------------------------------------
2# Copyright (R) Microsoft Corporation. All rights reserved.
3# Licensed under the MIT License.
4# --------------------------------------------------------------------------
5import nvtx
6from cuda import cudart
7
8
9class NvtxHelper:
10 def __init__(self, stages):
11 self.stages = stages
12 self.events = {}
13 for stage in stages:
14 for marker in ["start", "stop"]:
15 self.events[stage + "-" + marker] = cudart.cudaEventCreate()[1]
16 self.markers = {}
17
18 def start_profile(self, stage, color="blue"):
19 self.markers[stage] = nvtx.start_range(message=stage, color=color)
20 event_name = stage + "-start"
21 if event_name in self.events:
22 cudart.cudaEventRecord(self.events[event_name], 0)
23
24 def stop_profile(self, stage):
25 event_name = stage + "-stop"
26 if event_name in self.events:
27 cudart.cudaEventRecord(self.events[event_name], 0)
28 nvtx.end_range(self.markers[stage])
29
30 def print_latency(self):
31 for stage in self.stages:
32 latency = cudart.cudaEventElapsedTime(self.events[f"{stage}-start"], self.events[f"{stage}-stop"])[1]
33 print(f"{stage}: {latency:.2f} ms")
34 