[ROCm] Fix creation/destruction of HAL executable Make sure that if the creation fails and the object is left in an intermediate state it will be destroyed properly.
diff --git a/experimental/rocm/native_executable.c b/experimental/rocm/native_executable.c index 66fbb93..860cde0 100644 --- a/experimental/rocm/native_executable.c +++ b/experimental/rocm/native_executable.c
@@ -74,22 +74,30 @@ entry_count * sizeof(iree_hal_pipeline_layout_t*); iree_status_t status = iree_allocator_malloc(context->host_allocator, total_size, (void**)&executable); - executable->pipeline_layouts = - (void*)((char*)executable + sizeof(*executable) + - entry_count * sizeof(iree_hal_rocm_native_executable_function_t)); - hipModule_t module = NULL; if (iree_status_is_ok(status)) { + iree_hal_resource_initialize(&iree_hal_rocm_native_executable_vtable, + &executable->resource); + executable->context = context; status = ROCM_RESULT_TO_STATUS( - context->syms, hipModuleLoadDataEx(&module, hsaco_image, 0, NULL, NULL), + context->syms, + hipModuleLoadDataEx(&executable->module, hsaco_image, 0, NULL, NULL), "hipModuleLoadDataEx"); } + if (iree_status_is_ok(status)) { + executable->pipeline_layouts = + (void*)((char*)executable + sizeof(*executable) + + entry_count * + sizeof(iree_hal_rocm_native_executable_function_t)); + } + for (iree_host_size_t i = 0; i < entry_count; i++) { if (iree_status_is_ok(status)) { hipFunction_t function = NULL; const char* entry_name = flatbuffers_string_vec_at(entry_points_vec, i); status = ROCM_RESULT_TO_STATUS( - context->syms, hipModuleGetFunction(&function, module, entry_name), + context->syms, + hipModuleGetFunction(&function, executable->module, entry_name), "hipModuleGetFunction"); executable->entry_functions[i].rocm_function = function; executable->entry_functions[i].block_size_x = block_sizes_vec[i].x; @@ -101,13 +109,11 @@ } if (iree_status_is_ok(status)) { - iree_hal_resource_initialize(&iree_hal_rocm_native_executable_vtable, - &executable->resource); - executable->module = module; - executable->context = context; *out_executable = (iree_hal_executable_t*)executable; } else { - iree_hal_executable_destroy((iree_hal_executable_t*)executable); + if (executable) { + iree_hal_executable_destroy((iree_hal_executable_t*)executable); + } } IREE_TRACE_ZONE_END(z0); @@ -146,9 +152,25 @@ iree_allocator_t host_allocator = executable->context->host_allocator; IREE_TRACE_ZONE_BEGIN(z0); - for (iree_host_size_t i = 0; i < executable->entry_count; ++i) { - iree_hal_pipeline_layout_release(executable->pipeline_layouts[i]); + if (executable->module) { + iree_status_t status = ROCM_RESULT_TO_STATUS( + executable->context->syms, hipModuleUnload(executable->module), + "hipModuleUnload"); + if (!iree_status_is_ok(status)) { + fprintf(stderr, "Failed unloading ROCm module: "); + iree_status_fprint(stderr, status); + iree_status_free(status); + } } + + if (executable->pipeline_layouts) { + for (iree_host_size_t i = 0; i < executable->entry_count; ++i) { + if (executable->pipeline_layouts[i]) { + iree_hal_pipeline_layout_release(executable->pipeline_layouts[i]); + } + } + } + iree_allocator_free(host_allocator, executable); IREE_TRACE_ZONE_END(z0);