[metal] Unify pipeline object creation in MTLLibrary and source paths MTLFunction and MTLComputePipelineState creation are common between these two paths; so we can have one function for them.
diff --git a/experimental/metal/builtin_executables.m b/experimental/metal/builtin_executables.m index 696e1df..e20ff02 100644 --- a/experimental/metal/builtin_executables.m +++ b/experimental/metal/builtin_executables.m
@@ -77,9 +77,9 @@ id<MTLLibrary> library = nil; id<MTLFunction> function = nil; id<MTLComputePipelineState> pso = nil; - status = iree_hal_metal_compile_msl(iree_make_string_view(source_code.data, source_code.size), - IREE_SV(entry_point), device, compile_options, &library, - &function, &pso); + status = iree_hal_metal_compile_msl_and_create_pipeline_object( + iree_make_string_view(source_code.data, source_code.size), IREE_SV(entry_point), device, + compile_options, &library, &function, &pso); if (!iree_status_is_ok(status)) break; // Package required parameters for kernel launches for each entry point.
diff --git a/experimental/metal/metal_kernel_library.h b/experimental/metal/metal_kernel_library.h index 4d252f7..9cb11e7 100644 --- a/experimental/metal/metal_kernel_library.h +++ b/experimental/metal/metal_kernel_library.h
@@ -51,13 +51,11 @@ // Compiles the given |entry_point| in Metal |source_code| and writes the // |out_library|, |out_function|, and compute pipeline |out_pso| accordingly. -iree_status_t iree_hal_metal_compile_msl(iree_string_view_t source_code, - iree_string_view_t entry_point, - id<MTLDevice> device, - MTLCompileOptions* compile_options, - id<MTLLibrary>* out_library, - id<MTLFunction>* out_function, - id<MTLComputePipelineState>* out_pso); +iree_status_t iree_hal_metal_compile_msl_and_create_pipeline_object( + iree_string_view_t source_code, iree_string_view_t entry_point, + id<MTLDevice> device, MTLCompileOptions* compile_options, + id<MTLLibrary>* out_library, id<MTLFunction>* out_function, + id<MTLComputePipelineState>* out_pso); #ifdef __cplusplus } // extern "C"
diff --git a/experimental/metal/metal_kernel_library.m b/experimental/metal/metal_kernel_library.m index 44072d4..a31aad8 100644 --- a/experimental/metal/metal_kernel_library.m +++ b/experimental/metal/metal_kernel_library.m
@@ -145,14 +145,14 @@ (int)entry_point.size, entry_point.data); } +// Compiles the given |entry_point| in the MSL |source_code| into MTLLibrary and writes to +// |out_library|. The caller should release |out_library| after done. iree_status_t iree_hal_metal_compile_msl(iree_string_view_t source_code, iree_string_view_t entry_point, id<MTLDevice> device, MTLCompileOptions* compile_options, - id<MTLLibrary>* out_library, id<MTLFunction>* out_function, - id<MTLComputePipelineState>* out_pso) { + id<MTLLibrary>* out_library) { @autoreleasepool { NSError* error = nil; - NSString* shader_source = [[[NSString alloc] initWithBytes:source_code.data length:source_code.size @@ -160,75 +160,74 @@ *out_library = [device newLibraryWithSource:shader_source options:compile_options error:&error]; // +1 - if (*out_library == nil) { + if (IREE_UNLIKELY(*out_library == nil)) { return iree_hal_metal_get_invalid_kernel_status( "failed to create MTLLibrary from shader source", "when creating MTLLibrary with NSError: %.*s", error, entry_point, source_code.data); } + } + return iree_ok_status(); +} + +// Compiles the given |entry_point| in the MSL library |source_data| into MTLLibrary and writes to +// |out_library|. The caller should release |out_library| after done. +static iree_status_t iree_hal_metal_load_mtllib(iree_const_byte_span_t source_data, + iree_string_view_t entry_point, + id<MTLDevice> device, id<MTLLibrary>* out_library) { + @autoreleasepool { + NSError* error = nil; + dispatch_data_t data = dispatch_data_create(source_data.data, source_data.data_length, + /*queue=*/NULL, DISPATCH_DATA_DESTRUCTOR_DEFAULT); + *out_library = [device newLibraryWithData:data error:&error]; // +1 + if (IREE_UNLIKELY(*out_library == nil)) { + return iree_hal_metal_get_invalid_kernel_status( + "failed to create MTLLibrary from shader source", + "when creating MTLLibrary with NSError: %s", error, entry_point, NULL); + } + } + + return iree_ok_status(); +} + +// Creates MTL compute pipeline objects for the given |entry_point| in |library| and writes to +// |out_function| and |out_pso|. The caller should release |out_function| and |out_pso| after done. +static iree_status_t iree_hal_metal_create_pipline_object( + id<MTLLibrary> library, iree_string_view_t entry_point, const char* source_code, + id<MTLDevice> device, id<MTLFunction>* out_function, id<MTLComputePipelineState>* out_pso) { + @autoreleasepool { + NSError* error = nil; NSString* function_name = [[[NSString alloc] initWithBytes:entry_point.data length:entry_point.size encoding:[NSString defaultCStringEncoding]] autorelease]; - *out_function = [*out_library newFunctionWithName:function_name]; // +1 - if (*out_function == nil) { - [*out_library release]; + *out_function = [library newFunctionWithName:function_name]; // +1 + if (IREE_UNLIKELY(*out_function == nil)) { return iree_hal_metal_get_invalid_kernel_status("cannot find entry point in shader source", "when creating MTLFunction with NSError: %s", - error, entry_point, source_code.data); + error, entry_point, source_code); } // TODO(#14047): Enable async pipeline creation at runtime. *out_pso = [device newComputePipelineStateWithFunction:*out_function error:&error]; // +1 - if (*out_pso == nil) { + if (IREE_UNLIKELY(*out_pso == nil)) { [*out_function release]; - [*out_library release]; return iree_hal_metal_get_invalid_kernel_status( "invalid shader source", "when creating MTLComputePipelineState with NSError: %s", error, - entry_point, source_code.data); + entry_point, source_code); } } return iree_ok_status(); } -iree_status_t iree_hal_metal_load_mtllib(iree_const_byte_span_t source_data, - iree_string_view_t entry_point, id<MTLDevice> device, - id<MTLLibrary>* out_library, id<MTLFunction>* out_function, - id<MTLComputePipelineState>* out_pso) { - @autoreleasepool { - NSError* error = nil; - - dispatch_data_t data = dispatch_data_create(source_data.data, source_data.data_length, - /*queue=*/NULL, DISPATCH_DATA_DESTRUCTOR_DEFAULT); - *out_library = [device newLibraryWithData:data error:&error]; // +1 - if (*out_library == nil) { - return iree_hal_metal_get_invalid_kernel_status( - "failed to create MTLLibrary from shader source", - "when creating MTLLibrary with NSError: %s", error, entry_point, NULL); - } - - NSString* function_name = - [[[NSString alloc] initWithBytes:entry_point.data - length:entry_point.size - encoding:[NSString defaultCStringEncoding]] autorelease]; - *out_function = [*out_library newFunctionWithName:function_name]; // +1 - if (*out_function == nil) { - [*out_library release]; - return iree_hal_metal_get_invalid_kernel_status("cannot find entry point in shader source", - "when creating MTLFunction with NSError: %s", - error, entry_point, NULL); - } - - *out_pso = [device newComputePipelineStateWithFunction:*out_function error:&error]; // +1 - if (*out_pso == nil) { - [*out_function release]; - [*out_library release]; - return iree_hal_metal_get_invalid_kernel_status( - "invalid shader source", "when creating MTLComputePipelineState with NSError: %s", error, - entry_point, NULL); - } - } - return iree_ok_status(); +iree_status_t iree_hal_metal_compile_msl_and_create_pipeline_object( + iree_string_view_t source_code, iree_string_view_t entry_point, id<MTLDevice> device, + MTLCompileOptions* compile_options, id<MTLLibrary>* out_library, id<MTLFunction>* out_function, + id<MTLComputePipelineState>* out_pso) { + IREE_RETURN_IF_ERROR( + iree_hal_metal_compile_msl(source_code, entry_point, device, compile_options, out_library)); + return iree_hal_metal_create_pipline_object(*out_library, entry_point, source_code.data, device, + out_function, out_pso); } iree_status_t iree_hal_metal_kernel_library_create( @@ -296,22 +295,28 @@ id<MTLFunction> function = nil; id<MTLComputePipelineState> pso = nil; + flatbuffers_string_t source_code = NULL; flatbuffers_string_t entry_point = flatbuffers_string_vec_at(entry_points_vec, i); + iree_string_view_t entry_point_view = + iree_make_string_view(entry_point, flatbuffers_string_len(entry_point)); + if (shader_library_count != 0) { flatbuffers_string_t source_library = flatbuffers_string_vec_at(shader_libraries_vec, i); status = iree_hal_metal_load_mtllib( iree_make_const_byte_span(source_library, flatbuffers_string_len(source_library)), - iree_make_string_view(entry_point, flatbuffers_string_len(entry_point)), device, - &library, &function, &pso); + entry_point_view, device, &library); } else { - flatbuffers_string_t source_code = flatbuffers_string_vec_at(shader_sources_vec, i); + source_code = flatbuffers_string_vec_at(shader_sources_vec, i); status = iree_hal_metal_compile_msl( iree_make_string_view(source_code, flatbuffers_string_len(source_code)), - iree_make_string_view(entry_point, flatbuffers_string_len(entry_point)), device, - compile_options, &library, &function, &pso); + entry_point_view, device, compile_options, &library); } if (!iree_status_is_ok(status)) break; + status = iree_hal_metal_create_pipline_object(library, entry_point_view, source_code, device, + &function, &pso); + if (!iree_status_is_ok(status)) break; + // Package required parameters for kernel launches for each entry point. iree_hal_metal_kernel_params_t* params = &executable->entry_points[i]; params->library = library; @@ -351,9 +356,9 @@ for (iree_host_size_t i = 0; i < executable->entry_point_count; ++i) { iree_hal_metal_kernel_params_t* entry_point = &executable->entry_points[i]; - [entry_point->pso release]; - [entry_point->function release]; - [entry_point->library release]; + [entry_point->pso release]; // -1 + [entry_point->function release]; // -1 + [entry_point->library release]; // -1 iree_hal_pipeline_layout_release(entry_point->layout); } iree_allocator_free(executable->host_allocator, executable);