SWDEV-448586 - Added implementation for new API hipStreamBeginCaptureToGraph

Change-Id: I1ce802102cef2b66c92d3375f769983841de793f
This commit is contained in:
Anusha GodavarthySurya
2024-02-07 06:02:53 +00:00
parent 17d0c166d2
commit 4feb1f9337
7 changed files with 124 additions and 22 deletions
@@ -61,7 +61,7 @@
// - Reset any of the *_STEP_VERSION defines to zero if the corresponding *_MAJOR_VERSION increases
#define HIP_API_TABLE_STEP_VERSION 0
#define HIP_COMPILER_API_TABLE_STEP_VERSION 0
#define HIP_RUNTIME_API_TABLE_STEP_VERSION 1
#define HIP_RUNTIME_API_TABLE_STEP_VERSION 2
// HIP API interface
typedef hipError_t (*t___hipPopCallConfiguration)(dim3* gridDim, dim3* blockDim, size_t* sharedMem,
@@ -947,7 +947,11 @@ typedef hipError_t (*t_hipTexRefGetBorderColor)(float* pBorderColor,
typedef hipError_t (*t_hipTexRefGetArray)(hipArray_t* pArray, const textureReference* texRef);
typedef hipError_t (*t_hipGetProcAddress)(const char* symbol, void** pfn, int hipVersion, uint64_t flags,
hipDriverProcAddressQueryResult* symbolStatus);
typedef hipError_t (*t_hipStreamBeginCaptureToGraph)(hipStream_t stream, hipGraph_t graph,
const hipGraphNode_t* dependencies,
const hipGraphEdgeData* dependencyData,
size_t numDependencies,
hipStreamCaptureMode mode);
// HIP Compiler dispatch table
struct HipCompilerDispatchTable {
size_t size;
@@ -1412,4 +1416,5 @@ struct HipDispatchTable {
t_hipTexRefGetBorderColor hipTexRefGetBorderColor_fn;
t_hipTexRefGetArray hipTexRefGetArray_fn;
t_hipGetProcAddress hipGetProcAddress_fn;
t_hipStreamBeginCaptureToGraph hipStreamBeginCaptureToGraph_fn;
};
+40 -1
View File
@@ -408,7 +408,8 @@ enum hip_api_id_t {
HIP_API_ID_hipDrvGraphExecMemsetNodeSetParams = 388,
HIP_API_ID_hipTexRefGetArray = 389,
HIP_API_ID_hipTexRefGetBorderColor = 390,
HIP_API_ID_LAST = 390,
HIP_API_ID_hipStreamBeginCaptureToGraph = 391,
HIP_API_ID_LAST = 391,
HIP_API_ID_hipChooseDevice = HIP_API_ID_CONCAT(HIP_API_ID_,hipChooseDevice),
HIP_API_ID_hipGetDeviceProperties = HIP_API_ID_CONCAT(HIP_API_ID_,hipGetDeviceProperties),
@@ -795,6 +796,7 @@ static inline const char* hip_api_name(const uint32_t id) {
case HIP_API_ID_hipStreamAddCallback: return "hipStreamAddCallback";
case HIP_API_ID_hipStreamAttachMemAsync: return "hipStreamAttachMemAsync";
case HIP_API_ID_hipStreamBeginCapture: return "hipStreamBeginCapture";
case HIP_API_ID_hipStreamBeginCaptureToGraph: return "hipStreamBeginCaptureToGraph";
case HIP_API_ID_hipStreamCreate: return "hipStreamCreate";
case HIP_API_ID_hipStreamCreateWithFlags: return "hipStreamCreateWithFlags";
case HIP_API_ID_hipStreamCreateWithPriority: return "hipStreamCreateWithPriority";
@@ -1188,6 +1190,7 @@ static inline uint32_t hipApiIdByName(const char* name) {
if (strcmp("hipStreamAddCallback", name) == 0) return HIP_API_ID_hipStreamAddCallback;
if (strcmp("hipStreamAttachMemAsync", name) == 0) return HIP_API_ID_hipStreamAttachMemAsync;
if (strcmp("hipStreamBeginCapture", name) == 0) return HIP_API_ID_hipStreamBeginCapture;
if (strcmp("hipStreamBeginCaptureToGraph", name) == 0) return HIP_API_ID_hipStreamBeginCaptureToGraph;
if (strcmp("hipStreamCreate", name) == 0) return HIP_API_ID_hipStreamCreate;
if (strcmp("hipStreamCreateWithFlags", name) == 0) return HIP_API_ID_hipStreamCreateWithFlags;
if (strcmp("hipStreamCreateWithPriority", name) == 0) return HIP_API_ID_hipStreamCreateWithPriority;
@@ -3267,6 +3270,16 @@ typedef struct hip_api_data_s {
hipStream_t stream;
hipStreamCaptureMode mode;
} hipStreamBeginCapture;
struct {
hipStream_t stream;
hipGraph_t graph;
const hipGraphNode_t* dependencies;
hipGraphNode_t dependencies__val;
const hipGraphEdgeData* dependencyData;
hipGraphEdgeData dependencyData__val;
size_t numDependencies;
hipStreamCaptureMode mode;
} hipStreamBeginCaptureToGraph;
struct {
hipStream_t* stream;
hipStream_t stream__val;
@@ -5582,6 +5595,15 @@ typedef struct hip_api_data_s {
cb_data.args.hipStreamBeginCapture.stream = (hipStream_t)stream; \
cb_data.args.hipStreamBeginCapture.mode = (hipStreamCaptureMode)mode; \
};
// hipStreamBeginCaptureToGraph[('hipStream_t', 'stream'), ('hipGraph_t', 'graph'), ('const hipGraphNode_t*', 'dependencies'), ('const hipGraphEdgeData*', 'dependencyData'), ('size_t', 'numDependencies'), ('hipStreamCaptureMode', 'mode')]
#define INIT_hipStreamBeginCaptureToGraph_CB_ARGS_DATA(cb_data) { \
cb_data.args.hipStreamBeginCaptureToGraph.stream = (hipStream_t)stream; \
cb_data.args.hipStreamBeginCaptureToGraph.graph = (hipGraph_t)graph; \
cb_data.args.hipStreamBeginCaptureToGraph.dependencies = (const hipGraphNode_t*)dependencies; \
cb_data.args.hipStreamBeginCaptureToGraph.dependencyData = (const hipGraphEdgeData*)dependencyData; \
cb_data.args.hipStreamBeginCaptureToGraph.numDependencies = (size_t)numDependencies; \
cb_data.args.hipStreamBeginCaptureToGraph.mode = (hipStreamCaptureMode)mode; \
};
// hipStreamCreate[('hipStream_t*', 'stream')]
#define INIT_hipStreamCreate_CB_ARGS_DATA(cb_data) { \
cb_data.args.hipStreamCreate.stream = (hipStream_t*)stream; \
@@ -7231,6 +7253,11 @@ static inline void hipApiArgsInit(hip_api_id_t id, hip_api_data_t* data) {
// hipStreamBeginCapture[('hipStream_t', 'stream'), ('hipStreamCaptureMode', 'mode')]
case HIP_API_ID_hipStreamBeginCapture:
break;
// hipStreamBeginCaptureToGraph[('hipStream_t', 'stream'), ('hipGraph_t', 'graph'), ('const hipGraphNode_t*', 'dependencies'), ('const hipGraphEdgeData*', 'dependencyData'), ('size_t', 'numDependencies'), ('hipStreamCaptureMode', 'mode')]
case HIP_API_ID_hipStreamBeginCaptureToGraph:
if (data->args.hipStreamBeginCaptureToGraph.dependencies) data->args.hipStreamBeginCaptureToGraph.dependencies__val = *(data->args.hipStreamBeginCaptureToGraph.dependencies);
if (data->args.hipStreamBeginCaptureToGraph.dependencyData) data->args.hipStreamBeginCaptureToGraph.dependencyData__val = *(data->args.hipStreamBeginCaptureToGraph.dependencyData);
break;
// hipStreamCreate[('hipStream_t*', 'stream')]
case HIP_API_ID_hipStreamCreate:
if (data->args.hipStreamCreate.stream) data->args.hipStreamCreate.stream__val = *(data->args.hipStreamCreate.stream);
@@ -10156,6 +10183,18 @@ static inline const char* hipApiString(hip_api_id_t id, const hip_api_data_t* da
oss << ", mode="; roctracer::hip_support::detail::operator<<(oss, data->args.hipStreamBeginCapture.mode);
oss << ")";
break;
case HIP_API_ID_hipStreamBeginCaptureToGraph:
oss << "hipStreamBeginCaptureToGraph(";
oss << "stream="; roctracer::hip_support::detail::operator<<(oss, data->args.hipStreamBeginCaptureToGraph.stream);
oss << ", graph="; roctracer::hip_support::detail::operator<<(oss, data->args.hipStreamBeginCaptureToGraph.graph);
if (data->args.hipStreamBeginCaptureToGraph.dependencies == NULL) oss << ", dependencies=NULL";
else { oss << ", dependencies="; roctracer::hip_support::detail::operator<<(oss, data->args.hipStreamBeginCaptureToGraph.dependencies__val); }
if (data->args.hipStreamBeginCaptureToGraph.dependencyData == NULL) oss << ", dependencyData=NULL";
else { oss << ", dependencyData="; roctracer::hip_support::detail::operator<<(oss, data->args.hipStreamBeginCaptureToGraph.dependencyData__val); }
oss << ", numDependencies="; roctracer::hip_support::detail::operator<<(oss, data->args.hipStreamBeginCaptureToGraph.numDependencies);
oss << ", mode="; roctracer::hip_support::detail::operator<<(oss, data->args.hipStreamBeginCaptureToGraph.mode);
oss << ")";
break;
case HIP_API_ID_hipStreamCreate:
oss << "hipStreamCreate(";
if (data->args.hipStreamCreate.stream == NULL) oss << "stream=NULL";