Merge main -> google * 1b00f7c7e Limit the number of operands for FusionOfTensorOps to 32. (#6478) * 20ea69be9 Drop TESTONLY from sample target. (#6485) * 77a00779c Adding webgpu-headers submodule. (#6483) * e8b5f53fa Fixing task system off-by-one that was preventing proper distribution. (#6482) * e342abb22 Hiding tracing zones for imported VM functions. (#6481) * c31182585 Make tracing compatible with sanitizers. (#6479) * c519b6785 Disable VMA logging by default as it is slow and not generally useful. (#6480) * fc3ebe253 Interchange linalg generic op to move reduction loops to inner loops (#6470) * 565ba04d7 [CUDA] Fix workgroup memory creation for CUDA (#6455) * 9d75ceea4 Add an interface to `tensor::InsertSliceOp` to implement the `TiledOpInterface.. * 790ed1ad9 Bumping tracy submodule to 0.7.9+. (#6475) * 6c4cbb5e3 Adding iree_hal_command_buffer_push/pop_debug_group. (#6474) * 915df88ee Improving tracing of vulkan command buffers/dispatches. (#6473) * b3777ac8e ROCM HAL Maintenance (#6472) * 3135702b1 IREE Static Library Loader Runtime Demo (dylib-llvm-aot backend) (#6225) * 10460ea35 Add ir_to_markdown.py utility script. (#6467) * c821a6a54 Merge google -> main (#6468) PiperOrigin-RevId: 385672451
diff --git a/.gitmodules b/.gitmodules index a00f81a..2d8c87a 100644 --- a/.gitmodules +++ b/.gitmodules
@@ -47,3 +47,6 @@ [submodule "third_party/libyaml"] path = third_party/libyaml url = https://github.com/yaml/libyaml.git +[submodule "third_party/webgpu-headers"] + path = third_party/webgpu-headers + url = https://github.com/webgpu-native/webgpu-headers.git
diff --git a/SUBMODULE_VERSIONS.txt b/SUBMODULE_VERSIONS.txt index 5e2cfcd..5d408df 100644 --- a/SUBMODULE_VERSIONS.txt +++ b/SUBMODULE_VERSIONS.txt
@@ -11,6 +11,7 @@ f5417a4b6633c3217c9a1bc2f0c70b1454975ba7 third_party/spirv_headers b42009b3b9d4ca35bc703f5310eedc74f584be58 third_party/stblib 951452b6150ee86e2562aaf46a4bcceb4aa99697 third_party/tensorflow -f03b677ffa0fd96fcf859c32e79b740fac7dd59e third_party/tracy +50f7deb1a389bd3785c12fbe0be74128343f11f7 third_party/tracy 9d10a96f2d57c3c37e167f2e73c9a31ac2e51fa5 third_party/vulkan_headers 8d4a9e9174a9c6ad6a3a3ae981b915ef13fc12c4 third_party/vulkan_memory_allocator +c8e0b39f6f6f1edded5c4adf7d46aa4d2a95befe third_party/webgpu-headers
diff --git a/experimental/rocm/CMakeLists.txt b/experimental/rocm/CMakeLists.txt index db57bd3..185bd5c 100644 --- a/experimental/rocm/CMakeLists.txt +++ b/experimental/rocm/CMakeLists.txt
@@ -51,8 +51,6 @@ iree::base::internal iree::base::internal::flatcc iree::base::internal::synchronization - iree::base::logging - iree::base::status iree::base::tracing iree::hal iree::schemas::rocm_executable_def_c_fbs @@ -88,6 +86,7 @@ "dynamic_symbols_test.cc" DEPS ::dynamic_symbols + iree::base iree::testing::gtest iree::testing::gtest_main LABELS
diff --git a/experimental/rocm/api.h b/experimental/rocm/api.h index 638fedd..68fa191 100644 --- a/experimental/rocm/api.h +++ b/experimental/rocm/api.h
@@ -21,7 +21,7 @@ //===----------------------------------------------------------------------===// // ROCM driver creation options. -typedef struct { +typedef struct iree_hal_rocm_driver_options_t { // Index of the default ROCM device to use within the list of available // devices. int default_device_index;
diff --git a/experimental/rocm/context_wrapper.h b/experimental/rocm/context_wrapper.h index 514207d..2c51424 100644 --- a/experimental/rocm/context_wrapper.h +++ b/experimental/rocm/context_wrapper.h
@@ -13,7 +13,7 @@ // Structure to wrap all objects constant within a context. This makes it // simpler to pass it to the different objects and saves memory. -typedef struct { +typedef struct iree_hal_rocm_context_wrapper_t { hipCtx_t rocm_context; iree_allocator_t host_allocator; iree_hal_rocm_dynamic_symbols_t *syms;
diff --git a/experimental/rocm/descriptor_set_layout.c b/experimental/rocm/descriptor_set_layout.c index 0ca1cc8..69fb4af 100644 --- a/experimental/rocm/descriptor_set_layout.c +++ b/experimental/rocm/descriptor_set_layout.c
@@ -6,10 +6,12 @@ #include "experimental/rocm/descriptor_set_layout.h" -#include "experimental/rocm/status_util.h" +#include <stddef.h> + +#include "iree/base/api.h" #include "iree/base/tracing.h" -typedef struct { +typedef struct iree_hal_rocm_descriptor_set_layout_t { iree_hal_resource_t resource; iree_hal_rocm_context_wrapper_t *context; } iree_hal_rocm_descriptor_set_layout_t;
diff --git a/experimental/rocm/descriptor_set_layout.h b/experimental/rocm/descriptor_set_layout.h index 933c9c6..fb07b76 100644 --- a/experimental/rocm/descriptor_set_layout.h +++ b/experimental/rocm/descriptor_set_layout.h
@@ -8,6 +8,7 @@ #define IREE_HAL_ROCM_DESCRIPTOR_SET_LAYOUT_H_ #include "experimental/rocm/context_wrapper.h" +#include "iree/base/api.h" #include "iree/hal/api.h" #ifdef __cplusplus
diff --git a/experimental/rocm/direct_command_buffer.c b/experimental/rocm/direct_command_buffer.c index 9d56001..3d8706a 100644 --- a/experimental/rocm/direct_command_buffer.c +++ b/experimental/rocm/direct_command_buffer.c
@@ -6,10 +6,15 @@ #include "experimental/rocm/direct_command_buffer.h" +#include <assert.h> +#include <stddef.h> +#include <stdint.h> + +#include "experimental/rocm/dynamic_symbols.h" #include "experimental/rocm/native_executable.h" #include "experimental/rocm/rocm_buffer.h" -#include "experimental/rocm/rocm_event.h" #include "experimental/rocm/status_util.h" +#include "iree/base/api.h" #include "iree/base/tracing.h" // Command buffer implementation that directly maps to rocm direct. @@ -18,13 +23,13 @@ typedef struct { iree_hal_resource_t resource; - iree_hal_rocm_context_wrapper_t *context; + iree_hal_rocm_context_wrapper_t* context; iree_hal_command_buffer_mode_t mode; iree_hal_command_category_t allowed_categories; iree_hal_queue_affinity_t queue_affinity; size_t total_size; // Keep track of the current set of kernel arguments. - void *current_descriptor[]; + void* current_descriptor[]; } iree_hal_rocm_direct_command_buffer_t; #define IREE_HAL_ROCM_MAX_BINDING_COUNT 64 @@ -32,29 +37,29 @@ extern const iree_hal_command_buffer_vtable_t iree_hal_rocm_direct_command_buffer_vtable; -static iree_hal_rocm_direct_command_buffer_t * +static iree_hal_rocm_direct_command_buffer_t* iree_hal_rocm_direct_command_buffer_cast( - iree_hal_command_buffer_t *base_value) { + iree_hal_command_buffer_t* base_value) { IREE_HAL_ASSERT_TYPE(base_value, &iree_hal_rocm_direct_command_buffer_vtable); - return (iree_hal_rocm_direct_command_buffer_t *)base_value; + return (iree_hal_rocm_direct_command_buffer_t*)base_value; } iree_status_t iree_hal_rocm_direct_command_buffer_allocate( - iree_hal_rocm_context_wrapper_t *context, + iree_hal_rocm_context_wrapper_t* context, iree_hal_command_buffer_mode_t mode, iree_hal_command_category_t command_categories, iree_hal_queue_affinity_t queue_affinity, - iree_hal_command_buffer_t **out_command_buffer) { + iree_hal_command_buffer_t** out_command_buffer) { IREE_ASSERT_ARGUMENT(context); IREE_ASSERT_ARGUMENT(out_command_buffer); IREE_TRACE_ZONE_BEGIN(z0); - iree_hal_rocm_direct_command_buffer_t *command_buffer = NULL; + iree_hal_rocm_direct_command_buffer_t* command_buffer = NULL; size_t total_size = sizeof(*command_buffer) + - IREE_HAL_ROCM_MAX_BINDING_COUNT * sizeof(void *) + + IREE_HAL_ROCM_MAX_BINDING_COUNT * sizeof(void*) + IREE_HAL_ROCM_MAX_BINDING_COUNT * sizeof(hipDeviceptr_t); iree_status_t status = iree_allocator_malloc( - context->host_allocator, total_size, (void **)&command_buffer); + context->host_allocator, total_size, (void**)&command_buffer); if (iree_status_is_ok(status)) { iree_hal_resource_initialize(&iree_hal_rocm_direct_command_buffer_vtable, &command_buffer->resource); @@ -62,15 +67,15 @@ command_buffer->mode = mode; command_buffer->allowed_categories = command_categories; command_buffer->queue_affinity = queue_affinity; - hipDeviceptr_t *device_ptrs = - (hipDeviceptr_t *)(command_buffer->current_descriptor + - IREE_HAL_ROCM_MAX_BINDING_COUNT); + hipDeviceptr_t* device_ptrs = + (hipDeviceptr_t*)(command_buffer->current_descriptor + + IREE_HAL_ROCM_MAX_BINDING_COUNT); for (size_t i = 0; i < IREE_HAL_ROCM_MAX_BINDING_COUNT; i++) { command_buffer->current_descriptor[i] = &device_ptrs[i]; } command_buffer->total_size = total_size; - *out_command_buffer = (iree_hal_command_buffer_t *)command_buffer; + *out_command_buffer = (iree_hal_command_buffer_t*)command_buffer; } IREE_TRACE_ZONE_END(z0); @@ -78,8 +83,8 @@ } static void iree_hal_rocm_direct_command_buffer_destroy( - iree_hal_command_buffer_t *base_command_buffer) { - iree_hal_rocm_direct_command_buffer_t *command_buffer = + iree_hal_command_buffer_t* base_command_buffer) { + iree_hal_rocm_direct_command_buffer_t* command_buffer = iree_hal_rocm_direct_command_buffer_cast(base_command_buffer); IREE_TRACE_ZONE_BEGIN(z0); @@ -89,91 +94,103 @@ } static iree_hal_command_buffer_mode_t iree_hal_rocm_direct_command_buffer_mode( - const iree_hal_command_buffer_t *base_command_buffer) { - const iree_hal_rocm_direct_command_buffer_t *command_buffer = - (const iree_hal_rocm_direct_command_buffer_t *)(base_command_buffer); + const iree_hal_command_buffer_t* base_command_buffer) { + const iree_hal_rocm_direct_command_buffer_t* command_buffer = + (const iree_hal_rocm_direct_command_buffer_t*)(base_command_buffer); return command_buffer->mode; } static iree_hal_command_category_t iree_hal_rocm_direct_command_buffer_allowed_categories( - const iree_hal_command_buffer_t *base_command_buffer) { - const iree_hal_rocm_direct_command_buffer_t *command_buffer = - (const iree_hal_rocm_direct_command_buffer_t *)(base_command_buffer); + const iree_hal_command_buffer_t* base_command_buffer) { + const iree_hal_rocm_direct_command_buffer_t* command_buffer = + (const iree_hal_rocm_direct_command_buffer_t*)(base_command_buffer); return command_buffer->allowed_categories; } static iree_status_t iree_hal_rocm_direct_command_buffer_begin( - iree_hal_command_buffer_t *base_command_buffer) { + iree_hal_command_buffer_t* base_command_buffer) { return iree_ok_status(); } static iree_status_t iree_hal_rocm_direct_command_buffer_end( - iree_hal_command_buffer_t *base_command_buffer) { + iree_hal_command_buffer_t* base_command_buffer) { return iree_ok_status(); } +static void iree_hal_rocm_direct_command_buffer_begin_debug_group( + iree_hal_command_buffer_t* base_command_buffer, iree_string_view_t label, + iree_hal_label_color_t label_color, + const iree_hal_label_location_t* location) { + // TODO(benvanik): tracy event stack. +} + +static void iree_hal_rocm_direct_command_buffer_end_debug_group( + iree_hal_command_buffer_t* base_command_buffer) { + // TODO(benvanik): tracy event stack. +} + static iree_status_t iree_hal_rocm_direct_command_buffer_execution_barrier( - iree_hal_command_buffer_t *base_command_buffer, + iree_hal_command_buffer_t* base_command_buffer, iree_hal_execution_stage_t source_stage_mask, iree_hal_execution_stage_t target_stage_mask, iree_hal_execution_barrier_flags_t flags, iree_host_size_t memory_barrier_count, - const iree_hal_memory_barrier_t *memory_barriers, + const iree_hal_memory_barrier_t* memory_barriers, iree_host_size_t buffer_barrier_count, - const iree_hal_buffer_barrier_t *buffer_barriers) { + const iree_hal_buffer_barrier_t* buffer_barriers) { // TODO: Implement barrier return iree_ok_status(); } static iree_status_t iree_hal_rocm_direct_command_buffer_signal_event( - iree_hal_command_buffer_t *base_command_buffer, iree_hal_event_t *event, + iree_hal_command_buffer_t* base_command_buffer, iree_hal_event_t* event, iree_hal_execution_stage_t source_stage_mask) { // TODO: Implement barrier return iree_ok_status(); } static iree_status_t iree_hal_rocm_direct_command_buffer_reset_event( - iree_hal_command_buffer_t *base_command_buffer, iree_hal_event_t *event, + iree_hal_command_buffer_t* base_command_buffer, iree_hal_event_t* event, iree_hal_execution_stage_t source_stage_mask) { // TODO: Implement barrier return iree_ok_status(); } static iree_status_t iree_hal_rocm_direct_command_buffer_wait_events( - iree_hal_command_buffer_t *base_command_buffer, - iree_host_size_t event_count, const iree_hal_event_t **events, + iree_hal_command_buffer_t* base_command_buffer, + iree_host_size_t event_count, const iree_hal_event_t** events, iree_hal_execution_stage_t source_stage_mask, iree_hal_execution_stage_t target_stage_mask, iree_host_size_t memory_barrier_count, - const iree_hal_memory_barrier_t *memory_barriers, + const iree_hal_memory_barrier_t* memory_barriers, iree_host_size_t buffer_barrier_count, - const iree_hal_buffer_barrier_t *buffer_barriers) { + const iree_hal_buffer_barrier_t* buffer_barriers) { // TODO: Implement barrier return iree_ok_status(); } static iree_status_t iree_hal_rocm_direct_command_buffer_discard_buffer( - iree_hal_command_buffer_t *base_command_buffer, iree_hal_buffer_t *buffer) { + iree_hal_command_buffer_t* base_command_buffer, iree_hal_buffer_t* buffer) { // nothing to do. return iree_ok_status(); } // Splats a pattern value of 1, 2, or 4 bytes out to a 4 byte value. -static uint32_t iree_hal_rocm_splat_pattern(const void *pattern, +static uint32_t iree_hal_rocm_splat_pattern(const void* pattern, size_t pattern_length) { switch (pattern_length) { case 1: { - uint32_t pattern_value = *(const uint8_t *)(pattern); + uint32_t pattern_value = *(const uint8_t*)(pattern); return (pattern_value << 24) | (pattern_value << 16) | (pattern_value << 8) | pattern_value; } case 2: { - uint32_t pattern_value = *(const uint16_t *)(pattern); + uint32_t pattern_value = *(const uint16_t*)(pattern); return (pattern_value << 16) | pattern_value; } case 4: { - uint32_t pattern_value = *(const uint32_t *)(pattern); + uint32_t pattern_value = *(const uint32_t*)(pattern); return pattern_value; } default: @@ -182,11 +199,11 @@ } static iree_status_t iree_hal_rocm_direct_command_buffer_fill_buffer( - iree_hal_command_buffer_t *base_command_buffer, - iree_hal_buffer_t *target_buffer, iree_device_size_t target_offset, - iree_device_size_t length, const void *pattern, + iree_hal_command_buffer_t* base_command_buffer, + iree_hal_buffer_t* target_buffer, iree_device_size_t target_offset, + iree_device_size_t length, const void* pattern, iree_host_size_t pattern_length) { - iree_hal_rocm_direct_command_buffer_t *command_buffer = + iree_hal_rocm_direct_command_buffer_t* command_buffer = iree_hal_rocm_direct_command_buffer_cast(base_command_buffer); hipDeviceptr_t target_device_buffer = iree_hal_rocm_buffer_device_pointer( @@ -205,19 +222,19 @@ } static iree_status_t iree_hal_rocm_direct_command_buffer_update_buffer( - iree_hal_command_buffer_t *base_command_buffer, const void *source_buffer, - iree_host_size_t source_offset, iree_hal_buffer_t *target_buffer, + iree_hal_command_buffer_t* base_command_buffer, const void* source_buffer, + iree_host_size_t source_offset, iree_hal_buffer_t* target_buffer, iree_device_size_t target_offset, iree_device_size_t length) { return iree_make_status(IREE_STATUS_UNIMPLEMENTED, "need rocm implementation"); } static iree_status_t iree_hal_rocm_direct_command_buffer_copy_buffer( - iree_hal_command_buffer_t *base_command_buffer, - iree_hal_buffer_t *source_buffer, iree_device_size_t source_offset, - iree_hal_buffer_t *target_buffer, iree_device_size_t target_offset, + iree_hal_command_buffer_t* base_command_buffer, + iree_hal_buffer_t* source_buffer, iree_device_size_t source_offset, + iree_hal_buffer_t* target_buffer, iree_device_size_t target_offset, iree_device_size_t length) { - iree_hal_rocm_direct_command_buffer_t *command_buffer = + iree_hal_rocm_direct_command_buffer_t* command_buffer = iree_hal_rocm_direct_command_buffer_cast(base_command_buffer); hipDeviceptr_t target_device_buffer = iree_hal_rocm_buffer_device_pointer( @@ -237,9 +254,9 @@ } static iree_status_t iree_hal_rocm_direct_command_buffer_push_constants( - iree_hal_command_buffer_t *base_command_buffer, - iree_hal_executable_layout_t *executable_layout, iree_host_size_t offset, - const void *values, iree_host_size_t values_length) { + iree_hal_command_buffer_t* base_command_buffer, + iree_hal_executable_layout_t* executable_layout, iree_host_size_t offset, + const void* values, iree_host_size_t values_length) { return iree_make_status(IREE_STATUS_UNIMPLEMENTED, "need rocm implementation"); } @@ -251,20 +268,20 @@ } iree_hal_rocm_binding_mapping_t; // Helper to sort the binding based on their binding index. -static int compare_binding_index(const void *a, const void *b) { +static int compare_binding_index(const void* a, const void* b) { const iree_hal_rocm_binding_mapping_t buffer_a = - *(const iree_hal_rocm_binding_mapping_t *)a; + *(const iree_hal_rocm_binding_mapping_t*)a; const iree_hal_rocm_binding_mapping_t buffer_b = - *(const iree_hal_rocm_binding_mapping_t *)b; + *(const iree_hal_rocm_binding_mapping_t*)b; return buffer_a.binding < buffer_b.binding ? -1 : 1; } static iree_status_t iree_hal_rocm_direct_command_buffer_push_descriptor_set( - iree_hal_command_buffer_t *base_command_buffer, - iree_hal_executable_layout_t *executable_layout, uint32_t set, + iree_hal_command_buffer_t* base_command_buffer, + iree_hal_executable_layout_t* executable_layout, uint32_t set, iree_host_size_t binding_count, - const iree_hal_descriptor_set_binding_t *bindings) { - iree_hal_rocm_direct_command_buffer_t *command_buffer = + const iree_hal_descriptor_set_binding_t* bindings) { + iree_hal_rocm_direct_command_buffer_t* command_buffer = iree_hal_rocm_direct_command_buffer_cast(base_command_buffer); // Convention with the compiler side. We map bindings to kernel argument. // We compact the bindings to get a dense set of arguments and keep them order @@ -286,26 +303,26 @@ iree_hal_rocm_buffer_device_pointer( iree_hal_buffer_allocated_buffer(binding.buffer)) + iree_hal_buffer_byte_offset(binding.buffer) + binding.offset; - *((hipDeviceptr_t *)command_buffer->current_descriptor[i]) = device_ptr; + *((hipDeviceptr_t*)command_buffer->current_descriptor[i]) = device_ptr; } return iree_ok_status(); } static iree_status_t iree_hal_rocm_direct_command_buffer_bind_descriptor_set( - iree_hal_command_buffer_t *base_command_buffer, - iree_hal_executable_layout_t *executable_layout, uint32_t set, - iree_hal_descriptor_set_t *descriptor_set, + iree_hal_command_buffer_t* base_command_buffer, + iree_hal_executable_layout_t* executable_layout, uint32_t set, + iree_hal_descriptor_set_t* descriptor_set, iree_host_size_t dynamic_offset_count, - const iree_device_size_t *dynamic_offsets) { + const iree_device_size_t* dynamic_offsets) { return iree_make_status(IREE_STATUS_UNIMPLEMENTED, "need rocm implementation"); } static iree_status_t iree_hal_rocm_direct_command_buffer_dispatch( - iree_hal_command_buffer_t *base_command_buffer, - iree_hal_executable_t *executable, int32_t entry_point, + iree_hal_command_buffer_t* base_command_buffer, + iree_hal_executable_t* executable, int32_t entry_point, uint32_t workgroup_x, uint32_t workgroup_y, uint32_t workgroup_z) { - iree_hal_rocm_direct_command_buffer_t *command_buffer = + iree_hal_rocm_direct_command_buffer_t* command_buffer = iree_hal_rocm_direct_command_buffer_cast(base_command_buffer); iree_hal_rocm_direct_command_buffer_cast(base_command_buffer); @@ -327,9 +344,9 @@ } static iree_status_t iree_hal_rocm_direct_command_buffer_dispatch_indirect( - iree_hal_command_buffer_t *base_command_buffer, - iree_hal_executable_t *executable, int32_t entry_point, - iree_hal_buffer_t *workgroups_buffer, + iree_hal_command_buffer_t* base_command_buffer, + iree_hal_executable_t* executable, int32_t entry_point, + iree_hal_buffer_t* workgroups_buffer, iree_device_size_t workgroups_offset) { return iree_make_status(IREE_STATUS_UNIMPLEMENTED, "need rocm implementation"); @@ -343,6 +360,9 @@ iree_hal_rocm_direct_command_buffer_allowed_categories, .begin = iree_hal_rocm_direct_command_buffer_begin, .end = iree_hal_rocm_direct_command_buffer_end, + .begin_debug_group = + iree_hal_rocm_direct_command_buffer_begin_debug_group, + .end_debug_group = iree_hal_rocm_direct_command_buffer_end_debug_group, .execution_barrier = iree_hal_rocm_direct_command_buffer_execution_barrier, .signal_event = iree_hal_rocm_direct_command_buffer_signal_event,
diff --git a/experimental/rocm/direct_command_buffer.h b/experimental/rocm/direct_command_buffer.h index 70519c9..0145b15 100644 --- a/experimental/rocm/direct_command_buffer.h +++ b/experimental/rocm/direct_command_buffer.h
@@ -10,6 +10,7 @@ #include "experimental/rocm/context_wrapper.h" #include "experimental/rocm/dynamic_symbols.h" #include "experimental/rocm/rocm_headers.h" +#include "iree/base/api.h" #include "iree/hal/api.h" #ifdef __cplusplus
diff --git a/experimental/rocm/dynamic_symbols.c b/experimental/rocm/dynamic_symbols.c index 0643cf1..d676838 100644 --- a/experimental/rocm/dynamic_symbols.c +++ b/experimental/rocm/dynamic_symbols.c
@@ -6,7 +6,7 @@ #include "experimental/rocm/dynamic_symbols.h" -#include <stddef.h> +#include <string.h> #include "iree/base/internal/dynamic_library.h" #include "iree/base/target_platform.h" @@ -29,7 +29,7 @@ syms->loader_library, kName, (void **)&syms->rocmSymbolName)); \ } #define RC_PFN_STR_DECL(rocmSymbolName, ...) RC_PFN_DECL(rocmSymbolName, ...) -#include "experimental/rocm/dynamic_symbol_tables.h" +#include "experimental/rocm/dynamic_symbol_tables.h" // IWYU pragma: keep #undef RC_PFN_DECL #undef RC_PFN_STR_DECL return iree_ok_status();
diff --git a/experimental/rocm/dynamic_symbols.h b/experimental/rocm/dynamic_symbols.h index 2894626..00a58b0 100644 --- a/experimental/rocm/dynamic_symbols.h +++ b/experimental/rocm/dynamic_symbols.h
@@ -19,14 +19,14 @@ // loads all the function declared in `dynamic_symbol_tables.def` and fail if // any of the symbol is not available. The functions signatures are matching // the declarations in `hipruntime.h`. -typedef struct { +typedef struct iree_hal_rocm_dynamic_symbols_t { iree_dynamic_library_t *loader_library; #define RC_PFN_DECL(rocmSymbolName, ...) \ hipError_t (*rocmSymbolName)(__VA_ARGS__); #define RC_PFN_STR_DECL(rocmSymbolName, ...) \ const char *(*rocmSymbolName)(__VA_ARGS__); -#include "experimental/rocm/dynamic_symbol_tables.h" +#include "experimental/rocm/dynamic_symbol_tables.h" // IWYU pragma: export #undef RC_PFN_DECL #undef RC_PFN_STR_DECL } iree_hal_rocm_dynamic_symbols_t;
diff --git a/experimental/rocm/dynamic_symbols_test.cc b/experimental/rocm/dynamic_symbols_test.cc index dd33e07..06ea00a 100644 --- a/experimental/rocm/dynamic_symbols_test.cc +++ b/experimental/rocm/dynamic_symbols_test.cc
@@ -6,8 +6,10 @@ #include "experimental/rocm/dynamic_symbols.h" +#include <iostream> + +#include "iree/base/api.h" #include "iree/testing/gtest.h" -#include "iree/testing/status_matchers.h" namespace iree { namespace hal { @@ -25,7 +27,7 @@ iree_status_t status = iree_hal_rocm_dynamic_symbols_initialize( iree_allocator_system(), &symbols); if (!iree_status_is_ok(status)) { - IREE_LOG(WARNING) << "Symbols cannot be loaded, skipping test."; + std::cerr << "Symbols cannot be loaded, skipping test."; GTEST_SKIP(); }
diff --git a/experimental/rocm/event_semaphore.c b/experimental/rocm/event_semaphore.c index b05bafe..7fe6860 100644 --- a/experimental/rocm/event_semaphore.c +++ b/experimental/rocm/event_semaphore.c
@@ -6,10 +6,12 @@ #include "experimental/rocm/event_semaphore.h" -#include "experimental/rocm/status_util.h" +#include <stddef.h> + +#include "iree/base/api.h" #include "iree/base/tracing.h" -typedef struct { +typedef struct iree_hal_rocm_semaphore_t { iree_hal_resource_t resource; iree_hal_rocm_context_wrapper_t *context; uint64_t initial_value;
diff --git a/experimental/rocm/event_semaphore.h b/experimental/rocm/event_semaphore.h index 87ca683..5f32492 100644 --- a/experimental/rocm/event_semaphore.h +++ b/experimental/rocm/event_semaphore.h
@@ -7,8 +7,11 @@ #ifndef IREE_HAL_ROCM_SEMAPHORE_H_ #define IREE_HAL_ROCM_SEMAPHORE_H_ +#include <stdint.h> + #include "experimental/rocm/context_wrapper.h" #include "experimental/rocm/status_util.h" +#include "iree/base/api.h" #include "iree/hal/api.h" #ifdef __cplusplus
diff --git a/experimental/rocm/executable_layout.c b/experimental/rocm/executable_layout.c index ba88d32..6960cd2 100644 --- a/experimental/rocm/executable_layout.c +++ b/experimental/rocm/executable_layout.c
@@ -6,10 +6,12 @@ #include "experimental/rocm/executable_layout.h" -#include "experimental/rocm/status_util.h" +#include <stddef.h> + +#include "iree/base/api.h" #include "iree/base/tracing.h" -typedef struct { +typedef struct iree_hal_rocm_executable_layout_t { iree_hal_resource_t resource; iree_hal_rocm_context_wrapper_t *context; iree_host_size_t set_layout_count;
diff --git a/experimental/rocm/executable_layout.h b/experimental/rocm/executable_layout.h index c0469aa..cd0fccf 100644 --- a/experimental/rocm/executable_layout.h +++ b/experimental/rocm/executable_layout.h
@@ -8,6 +8,7 @@ #define IREE_HAL_ROCM_EXECUTABLE_LAYOUT_H_ #include "experimental/rocm/context_wrapper.h" +#include "iree/base/api.h" #include "iree/hal/api.h" #ifdef __cplusplus
diff --git a/experimental/rocm/native_executable.c b/experimental/rocm/native_executable.c index 0f1edf2..f72bd4f 100644 --- a/experimental/rocm/native_executable.c +++ b/experimental/rocm/native_executable.c
@@ -6,7 +6,11 @@ #include "experimental/rocm/native_executable.h" +#include <stddef.h> + +#include "experimental/rocm/dynamic_symbols.h" #include "experimental/rocm/status_util.h" +#include "iree/base/api.h" #include "iree/base/tracing.h" // flatcc schemas: @@ -14,14 +18,14 @@ #include "iree/schemas/rocm_executable_def_reader.h" #include "iree/schemas/rocm_executable_def_verifier.h" -typedef struct { +typedef struct iree_hal_rocm_native_executable_function_t { hipFunction_t rocm_function; uint32_t block_size_x; uint32_t block_size_y; uint32_t block_size_z; } iree_hal_rocm_native_executable_function_t; -typedef struct { +typedef struct iree_hal_rocm_native_executable_t { iree_hal_resource_t resource; iree_hal_rocm_context_wrapper_t *context; iree_host_size_t entry_count;
diff --git a/experimental/rocm/native_executable.h b/experimental/rocm/native_executable.h index 0b33b33..7a9229a 100644 --- a/experimental/rocm/native_executable.h +++ b/experimental/rocm/native_executable.h
@@ -7,8 +7,11 @@ #ifndef IREE_HAL_ROCM_NATIVE_EXECUTABLE_H_ #define IREE_HAL_ROCM_NATIVE_EXECUTABLE_H_ +#include <stdint.h> + #include "experimental/rocm/context_wrapper.h" #include "experimental/rocm/rocm_headers.h" +#include "iree/base/api.h" #include "iree/hal/api.h" #ifdef __cplusplus
diff --git a/experimental/rocm/nop_executable_cache.c b/experimental/rocm/nop_executable_cache.c index 35ec16c..6592727 100644 --- a/experimental/rocm/nop_executable_cache.c +++ b/experimental/rocm/nop_executable_cache.c
@@ -6,10 +6,14 @@ #include "experimental/rocm/nop_executable_cache.h" +#include <stdbool.h> +#include <stddef.h> + #include "experimental/rocm/native_executable.h" +#include "iree/base/api.h" #include "iree/base/tracing.h" -typedef struct { +typedef struct iree_hal_rocm_nop_executable_cache_t { iree_hal_resource_t resource; iree_hal_rocm_context_wrapper_t *context; } iree_hal_rocm_nop_executable_cache_t;
diff --git a/experimental/rocm/nop_executable_cache.h b/experimental/rocm/nop_executable_cache.h index 1262fea..d1b2fc1 100644 --- a/experimental/rocm/nop_executable_cache.h +++ b/experimental/rocm/nop_executable_cache.h
@@ -8,6 +8,7 @@ #define IREE_HAL_ROCM_NOP_EXECUTABLE_CACHE_H_ #include "experimental/rocm/context_wrapper.h" +#include "iree/base/api.h" #include "iree/hal/api.h" #ifdef __cplusplus
diff --git a/experimental/rocm/registration/CMakeLists.txt b/experimental/rocm/registration/CMakeLists.txt index f184333..36042bc 100644 --- a/experimental/rocm/registration/CMakeLists.txt +++ b/experimental/rocm/registration/CMakeLists.txt
@@ -16,8 +16,9 @@ SRCS "driver_module.c" DEPS + iree::base + iree::base::cc iree::base::core_headers - iree::base::status iree::base::tracing iree::hal experimental::rocm
diff --git a/experimental/rocm/registration/driver_module.c b/experimental/rocm/registration/driver_module.c index a9df558..2db41b4 100644 --- a/experimental/rocm/registration/driver_module.c +++ b/experimental/rocm/registration/driver_module.c
@@ -7,9 +7,10 @@ #include "experimental/rocm/registration/driver_module.h" #include <inttypes.h> +#include <stddef.h> #include "experimental/rocm/api.h" -#include "iree/base/target_platform.h" +#include "iree/base/api.h" #include "iree/base/tracing.h" #define IREE_HAL_ROCM_DRIVER_ID 0x524f434d0au // ROCM
diff --git a/experimental/rocm/registration/driver_module.h b/experimental/rocm/registration/driver_module.h index 6c55fea..efc1a70 100644 --- a/experimental/rocm/registration/driver_module.h +++ b/experimental/rocm/registration/driver_module.h
@@ -7,6 +7,7 @@ #ifndef IREE_HAL_ROCM_REGISTRATION_DRIVER_MODULE_H_ #define IREE_HAL_ROCM_REGISTRATION_DRIVER_MODULE_H_ +#include "iree/base/api.h" #include "iree/hal/api.h" #ifdef __cplusplus
diff --git a/experimental/rocm/rocm_allocator.c b/experimental/rocm/rocm_allocator.c index c30f4f3..cccdc9f 100644 --- a/experimental/rocm/rocm_allocator.c +++ b/experimental/rocm/rocm_allocator.c
@@ -6,11 +6,15 @@ #include "experimental/rocm/rocm_allocator.h" +#include <stddef.h> + +#include "experimental/rocm/dynamic_symbols.h" #include "experimental/rocm/rocm_buffer.h" #include "experimental/rocm/status_util.h" +#include "iree/base/api.h" #include "iree/base/tracing.h" -typedef struct iree_hal_rocm_allocator_s { +typedef struct iree_hal_rocm_allocator_t { iree_hal_resource_t resource; iree_hal_rocm_context_wrapper_t *context; } iree_hal_rocm_allocator_t;
diff --git a/experimental/rocm/rocm_allocator.h b/experimental/rocm/rocm_allocator.h index 1add674..a47480c 100644 --- a/experimental/rocm/rocm_allocator.h +++ b/experimental/rocm/rocm_allocator.h
@@ -9,6 +9,7 @@ #include "experimental/rocm/context_wrapper.h" #include "experimental/rocm/status_util.h" +#include "iree/base/api.h" #include "iree/hal/api.h" #ifdef __cplusplus
diff --git a/experimental/rocm/rocm_buffer.c b/experimental/rocm/rocm_buffer.c index 5269b4f..195e897 100644 --- a/experimental/rocm/rocm_buffer.c +++ b/experimental/rocm/rocm_buffer.c
@@ -6,11 +6,15 @@ #include "experimental/rocm/rocm_buffer.h" +#include <stddef.h> +#include <stdint.h> +#include <string.h> + #include "experimental/rocm/rocm_allocator.h" -#include "experimental/rocm/status_util.h" +#include "iree/base/api.h" #include "iree/base/tracing.h" -typedef struct iree_hal_rocm_buffer_s { +typedef struct iree_hal_rocm_buffer_t { iree_hal_buffer_t base; void *host_ptr; hipDeviceptr_t device_ptr;
diff --git a/experimental/rocm/rocm_buffer.h b/experimental/rocm/rocm_buffer.h index 7e24bc7..e898630 100644 --- a/experimental/rocm/rocm_buffer.h +++ b/experimental/rocm/rocm_buffer.h
@@ -8,6 +8,7 @@ #define IREE_HAL_ROCM_BUFFER_H_ #include "experimental/rocm/rocm_headers.h" +#include "iree/base/api.h" #include "iree/hal/api.h" #ifdef __cplusplus
diff --git a/experimental/rocm/rocm_device.c b/experimental/rocm/rocm_device.c index b84d2bc..67418f6 100644 --- a/experimental/rocm/rocm_device.c +++ b/experimental/rocm/rocm_device.c
@@ -6,7 +6,11 @@ #include "experimental/rocm/rocm_device.h" -#include "experimental/rocm/api.h" +#include <stddef.h> +#include <stdint.h> +#include <string.h> + +#include "experimental/rocm/context_wrapper.h" #include "experimental/rocm/descriptor_set_layout.h" #include "experimental/rocm/direct_command_buffer.h" #include "experimental/rocm/dynamic_symbols.h" @@ -22,7 +26,7 @@ // iree_hal_rocm_device_t //===----------------------------------------------------------------------===// -typedef struct { +typedef struct iree_hal_rocm_device_t { iree_hal_resource_t resource; iree_string_view_t identifier;
diff --git a/experimental/rocm/rocm_device.h b/experimental/rocm/rocm_device.h index 545436d..67e2f9b 100644 --- a/experimental/rocm/rocm_device.h +++ b/experimental/rocm/rocm_device.h
@@ -9,6 +9,7 @@ #include "experimental/rocm/api.h" #include "experimental/rocm/dynamic_symbols.h" +#include "iree/base/api.h" #include "iree/hal/api.h" #ifdef __cplusplus
diff --git a/experimental/rocm/rocm_driver.c b/experimental/rocm/rocm_driver.c index 4bc3323..8bd8ae6 100644 --- a/experimental/rocm/rocm_driver.c +++ b/experimental/rocm/rocm_driver.c
@@ -4,13 +4,18 @@ // See https://llvm.org/LICENSE.txt for license information. // SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception +#include <stdint.h> +#include <string.h> + #include "experimental/rocm/api.h" #include "experimental/rocm/dynamic_symbols.h" #include "experimental/rocm/rocm_device.h" #include "experimental/rocm/status_util.h" +#include "iree/base/api.h" #include "iree/base/tracing.h" +#include "iree/hal/api.h" -typedef struct { +typedef struct iree_hal_rocm_driver_t { iree_hal_resource_t resource; iree_allocator_t host_allocator; // Identifier used for the driver in the IREE driver registry.
diff --git a/experimental/rocm/rocm_event.c b/experimental/rocm/rocm_event.c index 2d17caa..48b20c3 100644 --- a/experimental/rocm/rocm_event.c +++ b/experimental/rocm/rocm_event.c
@@ -6,11 +6,13 @@ #include "experimental/rocm/rocm_event.h" -#include "experimental/rocm/status_util.h" +#include <stddef.h> + +#include "iree/base/api.h" #include "iree/base/tracing.h" // Dummy events for now, don't do anything. -typedef struct { +typedef struct iree_hal_rocm_event_t { iree_hal_resource_t resource; iree_hal_rocm_context_wrapper_t *context_wrapper; } iree_hal_rocm_event_t;
diff --git a/experimental/rocm/rocm_event.h b/experimental/rocm/rocm_event.h index 6c1ece6..73041a4 100644 --- a/experimental/rocm/rocm_event.h +++ b/experimental/rocm/rocm_event.h
@@ -9,6 +9,7 @@ #include "experimental/rocm/context_wrapper.h" #include "experimental/rocm/rocm_headers.h" +#include "iree/base/api.h" #include "iree/hal/api.h" #ifdef __cplusplus
diff --git a/experimental/rocm/rocm_headers.h b/experimental/rocm/rocm_headers.h index 056e8c0..6f5da96 100644 --- a/experimental/rocm/rocm_headers.h +++ b/experimental/rocm/rocm_headers.h
@@ -11,6 +11,6 @@ #error 32-bit not supported on ROCm #endif // defined(IREE_PTR_SIZE_32) -#include "hip/hip_runtime.h" +#include "hip/hip_runtime.h" // IWYU pragma: export #endif // IREE_HAL_ROCM_ROCM_HEADERS_H_
diff --git a/experimental/rocm/status_util.c b/experimental/rocm/status_util.c index e1b6fd6..a084c3f 100644 --- a/experimental/rocm/status_util.c +++ b/experimental/rocm/status_util.c
@@ -6,6 +6,8 @@ #include "experimental/rocm/status_util.h" +#include <stddef.h> + #include "experimental/rocm/dynamic_symbols.h" iree_status_t iree_hal_rocm_result_to_status(
diff --git a/experimental/rocm/status_util.h b/experimental/rocm/status_util.h index 31e9404..b459591 100644 --- a/experimental/rocm/status_util.h +++ b/experimental/rocm/status_util.h
@@ -7,6 +7,8 @@ #ifndef IREE_HAL_ROCM_STATUS_UTIL_H_ #define IREE_HAL_ROCM_STATUS_UTIL_H_ +#include <stdint.h> + #include "experimental/rocm/dynamic_symbols.h" #include "iree/base/api.h"
diff --git a/iree/base/tracing.cc b/iree/base/tracing.cc index 917370d..870eabd 100644 --- a/iree/base/tracing.cc +++ b/iree/base/tracing.cc
@@ -204,8 +204,10 @@ } // extern "C" #endif // __cplusplus -#if defined(__cplusplus) && \ - (IREE_TRACING_FEATURES & IREE_TRACING_FEATURE_ALLOCATION_TRACKING) +#if defined(__cplusplus) && \ + (IREE_TRACING_FEATURES & IREE_TRACING_FEATURE_ALLOCATION_TRACKING) && \ + !IREE_SANITIZER_ADDRESS && !IREE_SANITIZER_MEMORY && \ + !IREE_SANITIZER_THREAD void* operator new(size_t count) noexcept { auto ptr = malloc(count);
diff --git a/iree/compiler/Codegen/LLVMGPU/ConvertToLLVM.cpp b/iree/compiler/Codegen/LLVMGPU/ConvertToLLVM.cpp index 297b325..28a4dbc0 100644 --- a/iree/compiler/Codegen/LLVMGPU/ConvertToLLVM.cpp +++ b/iree/compiler/Codegen/LLVMGPU/ConvertToLLVM.cpp
@@ -62,15 +62,49 @@ } }; -/// Pass to test scalarization pattern. -class TestLLVMGPUScalarizeMathOpPass - : public TestLLVMGPUScalarizeMathOpBase<TestLLVMGPUScalarizeMathOpPass> { +struct ConvertSharedMemAllocOp : public OpRewritePattern<memref::AllocOp> { + using OpRewritePattern<memref::AllocOp>::OpRewritePattern; + + LogicalResult matchAndRewrite(memref::AllocOp allocOp, + PatternRewriter &rewriter) const override { + if (allocOp.getType().getMemorySpaceAsInt() != 3) return failure(); + ArrayRef<int64_t> shape = allocOp.getType().getShape(); + if (llvm::any_of( + shape, [](int64_t dim) { return dim == ShapedType::kDynamicSize; })) + return failure(); + // In CUDA workgroup memory is represented by a global variable. + MemRefType allocType = allocOp.getType(); + auto funcOp = allocOp->getParentOfType<FuncOp>(); + auto moduleOp = funcOp->getParentOfType<ModuleOp>(); + SymbolTable symbolTable(moduleOp); + OpBuilder::InsertionGuard guard(rewriter); + rewriter.setInsertionPoint(&moduleOp.front()); + auto global = rewriter.create<memref::GlobalOp>( + funcOp.getLoc(), "__shared_memory__", + /*sym_visibility=*/rewriter.getStringAttr("private"), + /*type=*/allocType, + /*initial_value=*/ElementsAttr(), + /*constant=*/false); + symbolTable.insert(global); + + rewriter.setInsertionPointToStart(&(*funcOp.getBody().begin())); + rewriter.replaceOpWithNewOp<memref::GetGlobalOp>(allocOp, global.type(), + global.getName()); + return success(); + } +}; + +/// Pass to test in dialect transformation used to legalize the IR before +/// convertToNVVM/ConvertToROCDL. +class TestLLVMGPULegalizeOpPass + : public TestLLVMGPUScalarizeMathOpBase<TestLLVMGPULegalizeOpPass> { void getDependentDialects(DialectRegistry ®istry) const override { registry.insert<vector::VectorDialect>(); } void runOnOperation() override { OwningRewritePatternList patterns(&getContext()); populateScalarizeMathOps(patterns); + populateConvertSharedMemoryAllocOps(patterns); (void)applyPatternsAndFoldGreedily(getOperation(), std::move(patterns)); } }; @@ -290,8 +324,12 @@ patterns.getContext()); } -std::unique_ptr<OperationPass<FuncOp>> createTestLLVMGPUScalarizeMathOpPass() { - return std::make_unique<TestLLVMGPUScalarizeMathOpPass>(); +void populateConvertSharedMemoryAllocOps(RewritePatternSet &patterns) { + patterns.add<ConvertSharedMemAllocOp>(patterns.getContext()); +} + +std::unique_ptr<OperationPass<ModuleOp>> createTestLLVMGPULegalizePass() { + return std::make_unique<TestLLVMGPULegalizeOpPass>(); } } // namespace iree_compiler
diff --git a/iree/compiler/Codegen/LLVMGPU/ConvertToLLVM.h b/iree/compiler/Codegen/LLVMGPU/ConvertToLLVM.h index f639bab..62ed632 100644 --- a/iree/compiler/Codegen/LLVMGPU/ConvertToLLVM.h +++ b/iree/compiler/Codegen/LLVMGPU/ConvertToLLVM.h
@@ -18,6 +18,9 @@ void populateScalarizeMathOps(RewritePatternSet &patterns); +/// Add patterns to convert AllocOp of shared memory to a global variable. +void populateConvertSharedMemoryAllocOps(RewritePatternSet &patterns); + } // namespace iree_compiler } // namespace mlir
diff --git a/iree/compiler/Codegen/LLVMGPU/ConvertToNVVM.cpp b/iree/compiler/Codegen/LLVMGPU/ConvertToNVVM.cpp index 90a9976..7544522 100644 --- a/iree/compiler/Codegen/LLVMGPU/ConvertToNVVM.cpp +++ b/iree/compiler/Codegen/LLVMGPU/ConvertToNVVM.cpp
@@ -52,6 +52,7 @@ { OwningRewritePatternList patterns(&getContext()); populateScalarizeMathOps(patterns); + populateConvertSharedMemoryAllocOps(patterns); vector::populateVectorToVectorCanonicalizationPatterns(patterns); vector::populateVectorContractLoweringPatterns( patterns,
diff --git a/iree/compiler/Codegen/LLVMGPU/ConvertToROCDL.cpp b/iree/compiler/Codegen/LLVMGPU/ConvertToROCDL.cpp index 92cf4d2..b6cf257 100644 --- a/iree/compiler/Codegen/LLVMGPU/ConvertToROCDL.cpp +++ b/iree/compiler/Codegen/LLVMGPU/ConvertToROCDL.cpp
@@ -52,6 +52,7 @@ { OwningRewritePatternList patterns(&getContext()); populateScalarizeMathOps(patterns); + populateConvertSharedMemoryAllocOps(patterns); vector::populateVectorToVectorCanonicalizationPatterns(patterns); vector::populateVectorContractLoweringPatterns( patterns,
diff --git a/iree/compiler/Codegen/LLVMGPU/test/BUILD b/iree/compiler/Codegen/LLVMGPU/test/BUILD index 831c607..e291d59 100644 --- a/iree/compiler/Codegen/LLVMGPU/test/BUILD +++ b/iree/compiler/Codegen/LLVMGPU/test/BUILD
@@ -25,7 +25,7 @@ "nvvm_pipeline_test.mlir", "remove_loops.mlir", "rocdl_pipeline_test.mlir", - "scalarize.mlir", + "legalize.mlir", "vectorization.mlir", ], include = ["*.mlir"],
diff --git a/iree/compiler/Codegen/LLVMGPU/test/CMakeLists.txt b/iree/compiler/Codegen/LLVMGPU/test/CMakeLists.txt index b18ba76..51cb48c 100644 --- a/iree/compiler/Codegen/LLVMGPU/test/CMakeLists.txt +++ b/iree/compiler/Codegen/LLVMGPU/test/CMakeLists.txt
@@ -17,10 +17,10 @@ "convert_to_nvvm.mlir" "convert_to_rocdl.mlir" "distribute_to_thread.mlir" + "legalize.mlir" "nvvm_pipeline_test.mlir" "remove_loops.mlir" "rocdl_pipeline_test.mlir" - "scalarize.mlir" "vectorization.mlir" DATA iree::tools::IreeFileCheck
diff --git a/iree/compiler/Codegen/LLVMGPU/test/scalarize.mlir b/iree/compiler/Codegen/LLVMGPU/test/legalize.mlir similarity index 79% rename from iree/compiler/Codegen/LLVMGPU/test/scalarize.mlir rename to iree/compiler/Codegen/LLVMGPU/test/legalize.mlir index da17849..d2df6b9 100644 --- a/iree/compiler/Codegen/LLVMGPU/test/scalarize.mlir +++ b/iree/compiler/Codegen/LLVMGPU/test/legalize.mlir
@@ -1,4 +1,4 @@ -// RUN: iree-opt -iree-test-llvmgpu-scalarize-math-op %s | IreeFileCheck %s +// RUN: iree-opt -iree-test-llvmgpu-legalize-ops -split-input-file %s | IreeFileCheck %s // CHECK-LABEL: func @scalarize func @scalarize( @@ -34,3 +34,16 @@ %1 = math.powf %arg1, %arg2 : vector<2xf32> return %0, %1 : vector<3x1x2xf32>, vector<2xf32> } + +// ----- + +// CHECK: memref.global "private" @__shared_memory__ : memref<16x16xf32, 3> +// CHECK: func @allocation +// CHECK: %[[A:.*]] = memref.get_global @__shared_memory__ : memref<16x16xf32, 3> +// CHECK: memref.store %{{.*}}, %[[A]][%{{.*}}, %{{.*}}] : memref<16x16xf32, 3> +func @allocation(%arg0: f32) { + %0 = memref.alloc() : memref<16x16xf32, 3> + %c0 = constant 0 : index + memref.store %arg0, %0[%c0, %c0] : memref<16x16xf32, 3> + return +}
diff --git a/iree/compiler/Codegen/Passes.h b/iree/compiler/Codegen/Passes.h index e8c09a9..c42ce1c 100644 --- a/iree/compiler/Codegen/Passes.h +++ b/iree/compiler/Codegen/Passes.h
@@ -314,7 +314,7 @@ // Test passes //------------------------------------------------------------------------------ -std::unique_ptr<OperationPass<FuncOp>> createTestLLVMGPUScalarizeMathOpPass(); +std::unique_ptr<OperationPass<ModuleOp>> createTestLLVMGPULegalizePass(); } // namespace iree_compiler } // namespace mlir
diff --git a/iree/compiler/Codegen/Passes.td b/iree/compiler/Codegen/Passes.td index 8fa0e9e..3029395 100644 --- a/iree/compiler/Codegen/Passes.td +++ b/iree/compiler/Codegen/Passes.td
@@ -260,9 +260,9 @@ //------------------------------------------------------------------------------ def TestLLVMGPUScalarizeMathOp : - Pass<"iree-test-llvmgpu-scalarize-math-op", "FuncOp"> { - let summary = "Test pass for scalarization patterns."; - let constructor = "mlir::iree_compiler::createTestLLVMGPUScalarizeMathOpPass()"; + Pass<"iree-test-llvmgpu-legalize-ops", "ModuleOp"> { + let summary = "Test pass for several legalization patterns."; + let constructor = "mlir::iree_compiler::createTestLLVMGPULegalizePass()"; } #endif // IREE_DIALECT_FLOW_PASSES
diff --git a/iree/compiler/Dialect/Flow/Transforms/BUILD b/iree/compiler/Dialect/Flow/Transforms/BUILD index c521376..fcc530c 100644 --- a/iree/compiler/Dialect/Flow/Transforms/BUILD +++ b/iree/compiler/Dialect/Flow/Transforms/BUILD
@@ -43,6 +43,7 @@ "FusionOfTensorOps.cpp", "HoistUnstreamableOps.cpp", "InjectDispatchTracing.cpp", + "InterchangeGenericOp.cpp", "OutlineDispatchRegions.cpp", "OutlineLargeConstants.cpp", "PadLinalgOps.cpp",
diff --git a/iree/compiler/Dialect/Flow/Transforms/CMakeLists.txt b/iree/compiler/Dialect/Flow/Transforms/CMakeLists.txt index 33762e2..ba6b504 100644 --- a/iree/compiler/Dialect/Flow/Transforms/CMakeLists.txt +++ b/iree/compiler/Dialect/Flow/Transforms/CMakeLists.txt
@@ -40,6 +40,7 @@ "FusionOfTensorOps.cpp" "HoistUnstreamableOps.cpp" "InjectDispatchTracing.cpp" + "InterchangeGenericOp.cpp" "OutlineDispatchRegions.cpp" "OutlineLargeConstants.cpp" "PadLinalgOps.cpp"
diff --git a/iree/compiler/Dialect/Flow/Transforms/FusionOfTensorOps.cpp b/iree/compiler/Dialect/Flow/Transforms/FusionOfTensorOps.cpp index 9032285..6689134 100644 --- a/iree/compiler/Dialect/Flow/Transforms/FusionOfTensorOps.cpp +++ b/iree/compiler/Dialect/Flow/Transforms/FusionOfTensorOps.cpp
@@ -64,6 +64,14 @@ } } + // Limit the number of operands. We have hard limit (32) of bindings + // passing down to HAL. Set the number to be as same as the limit -- + // IREE_HAL_MODULE_MAX_DESCRIPTOR_BINDING_COUNT. + constexpr int64_t kIreeMaxOperandCount = 32; + auto numOperands = producer.getOwner()->getNumOperands() + + consumer.getOwner()->getNumOperands() - 1; + if (numOperands >= kIreeMaxOperandCount) return false; + llvm::SmallDenseSet<Operation *, 4> numUsers; for (Operation *user : producer.getUsers()) { if (isa<linalg::GenericOp>(user)) continue;
diff --git a/iree/compiler/Dialect/Flow/Transforms/InterchangeGenericOp.cpp b/iree/compiler/Dialect/Flow/Transforms/InterchangeGenericOp.cpp new file mode 100644 index 0000000..8748060 --- /dev/null +++ b/iree/compiler/Dialect/Flow/Transforms/InterchangeGenericOp.cpp
@@ -0,0 +1,78 @@ +// Copyright 2021 The IREE Authors +// +// Licensed under the Apache License v2.0 with LLVM Exceptions. +// See https://llvm.org/LICENSE.txt for license information. +// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception + +//===--------------- InterchangeGenericOp.cpp -----------------------------===// +// +// Interchange loops in generic ops to force the reduction loops to be the most +// inner loops. +// +//===----------------------------------------------------------------------===// + +#include "iree/compiler/Dialect/Flow/Transforms/PassDetail.h" +#include "iree/compiler/Dialect/Flow/Transforms/Passes.h" +#include "mlir/Dialect/Affine/IR/AffineOps.h" +#include "mlir/Dialect/Linalg/IR/LinalgOps.h" +#include "mlir/Dialect/Linalg/Transforms/Transforms.h" +#include "mlir/Transforms/GreedyPatternRewriteDriver.h" + +namespace mlir { +namespace iree_compiler { +namespace IREE { +namespace Flow { + +namespace { + +struct GenericOpInterchangePattern + : public OpRewritePattern<linalg::GenericOp> { + using OpRewritePattern<linalg::GenericOp>::OpRewritePattern; + LogicalResult matchAndRewrite(linalg::GenericOp genericOp, + PatternRewriter &rewriter) const override { + SmallVector<unsigned> interchange; + bool needInterchange = false; + unsigned numParallelLoop = genericOp.getNumParallelLoops(); + if (numParallelLoop == 0) return failure(); + for (auto iter : llvm::enumerate(genericOp.iterator_types())) { + if (isParallelIterator(iter.value())) { + interchange.push_back(iter.index()); + if (iter.index() >= numParallelLoop) needInterchange = true; + } + } + // If all the parallel loops are outter loops skip the pattern. + if (!needInterchange) return failure(); + for (auto iter : llvm::enumerate(genericOp.iterator_types())) { + if (isReductionIterator(iter.value())) + interchange.push_back(iter.index()); + } + rewriter.updateRootInPlace(genericOp, [&]() { + interchangeGenericOp(rewriter, genericOp, interchange); + }); + return success(); + } +}; + +struct InterchangeGenericOpsPass + : public InterchangeGenericBase<InterchangeGenericOpsPass> { + void getDependentDialects(DialectRegistry ®istry) const override { + registry.insert<linalg::LinalgDialect>(); + } + + void runOnOperation() override { + OwningRewritePatternList patterns(&getContext()); + patterns.add<GenericOpInterchangePattern>(&getContext()); + (void)applyPatternsAndFoldGreedily(getOperation(), std::move(patterns)); + } +}; + +} // namespace + +std::unique_ptr<OperationPass<FuncOp>> createInterchangeLinalgGenericPass() { + return std::make_unique<InterchangeGenericOpsPass>(); +} + +} // namespace Flow +} // namespace IREE +} // namespace iree_compiler +} // namespace mlir
diff --git a/iree/compiler/Dialect/Flow/Transforms/Passes.cpp b/iree/compiler/Dialect/Flow/Transforms/Passes.cpp index 6f2288f..2a5e18b 100644 --- a/iree/compiler/Dialect/Flow/Transforms/Passes.cpp +++ b/iree/compiler/Dialect/Flow/Transforms/Passes.cpp
@@ -105,6 +105,7 @@ passManager.addNestedPass<FuncOp>( mlir::createConvertElementwiseToLinalgPass()); passManager.addNestedPass<FuncOp>(mlir::createLinalgFoldUnitExtentDimsPass()); + passManager.addNestedPass<FuncOp>(createInterchangeLinalgGenericPass()); passManager.addNestedPass<FuncOp>(mlir::createCanonicalizerPass()); passManager.addNestedPass<FuncOp>(createFusionOfTensorOpsPass()); passManager.addNestedPass<FuncOp>(
diff --git a/iree/compiler/Dialect/Flow/Transforms/Passes.h b/iree/compiler/Dialect/Flow/Transforms/Passes.h index 2c1fbb5..2a58c1f 100644 --- a/iree/compiler/Dialect/Flow/Transforms/Passes.h +++ b/iree/compiler/Dialect/Flow/Transforms/Passes.h
@@ -57,6 +57,10 @@ /// Creates a pass to fuse Linalg operations on tensors. std::unique_ptr<Pass> createFusionOfTensorOpsPass(); +/// Create a pass to interchange generic ops to force the reduction loop to be +/// the most inner loops. +std::unique_ptr<OperationPass<FuncOp>> createInterchangeLinalgGenericPass(); + // Convert operations to equivalent flow.tensor.* ops. This is run after // dispatch region creation to catch operations that were left outside of // dispatch regions and could be represented as flow.tensor.* ops.
diff --git a/iree/compiler/Dialect/Flow/Transforms/Passes.td b/iree/compiler/Dialect/Flow/Transforms/Passes.td index 5788193..e15d56b 100644 --- a/iree/compiler/Dialect/Flow/Transforms/Passes.td +++ b/iree/compiler/Dialect/Flow/Transforms/Passes.td
@@ -119,4 +119,10 @@ let constructor = "mlir::iree_compiler::IREE::Flow::createStripAndSplatConstantVariablesPass()"; } +def InterchangeGeneric : + Pass<"iree-flow-interchange-generic-op", "FuncOp"> { + let summary = "Interchange generic op loops to have all the reduction loops to be inner loops."; + let constructor = "mlir::iree_compiler::IREE::Flow::createInterchangeLinalgGenericPass()"; +} + #endif // IREE_DIALECT_FLOW_PASSES
diff --git a/iree/compiler/Dialect/Flow/Transforms/test/BUILD b/iree/compiler/Dialect/Flow/Transforms/test/BUILD index 9953942..d2a70f8 100644 --- a/iree/compiler/Dialect/Flow/Transforms/test/BUILD +++ b/iree/compiler/Dialect/Flow/Transforms/test/BUILD
@@ -29,6 +29,7 @@ "form_streams.mlir", "hoist_unstreamable_ops.mlir", "inject_dispatch_tracing.mlir", + "interchange.mlir", "outline_dispatch_regions.mlir", "outline_large_constants.mlir", "pad_linalg_ops.mlir",
diff --git a/iree/compiler/Dialect/Flow/Transforms/test/CMakeLists.txt b/iree/compiler/Dialect/Flow/Transforms/test/CMakeLists.txt index b43f4aa..21cdce8 100644 --- a/iree/compiler/Dialect/Flow/Transforms/test/CMakeLists.txt +++ b/iree/compiler/Dialect/Flow/Transforms/test/CMakeLists.txt
@@ -26,6 +26,7 @@ "form_streams.mlir" "hoist_unstreamable_ops.mlir" "inject_dispatch_tracing.mlir" + "interchange.mlir" "outline_dispatch_regions.mlir" "outline_large_constants.mlir" "pad_linalg_ops.mlir"
diff --git a/iree/compiler/Dialect/Flow/Transforms/test/interchange.mlir b/iree/compiler/Dialect/Flow/Transforms/test/interchange.mlir new file mode 100644 index 0000000..8c9a019 --- /dev/null +++ b/iree/compiler/Dialect/Flow/Transforms/test/interchange.mlir
@@ -0,0 +1,23 @@ +// RUN: iree-opt -iree-flow-interchange-generic-op %s | IreeFileCheck %s + +// CHECK: #[[MAP0:.+]] = affine_map<(d0, d1, d2, d3) -> (d0, d3, d2)> +// CHECK: #[[MAP1:.+]] = affine_map<(d0, d1, d2, d3) -> (d3, d0, d1)> +// CHECK: #[[MAP2:.+]] = affine_map<(d0, d1, d2, d3) -> (d2, d0, d1)> +// CHECK: func @interchange +// CHECK: linalg.generic {indexing_maps = [#[[MAP0]], #[[MAP1]], #[[MAP2]]] +// CHECK-SAME: iterator_types = ["parallel", "parallel", "parallel", "reduction"]} +func @interchange(%arg0: tensor<?x?x?xf32>, %arg1: tensor<?x?x?xf32>, %arg2: tensor<?x?x?xf32>) -> (tensor<?x?x?xf32>) { + %0 = linalg.generic {indexing_maps = [ + affine_map<(d0, d1, d2, d3) -> (d1, d0, d3)>, + affine_map<(d0, d1, d2, d3) -> (d0, d1, d2)>, + affine_map<(d0, d1, d2, d3) -> (d3, d1, d2)>], + iterator_types = ["reduction", "parallel", "parallel", "parallel"]} + ins(%arg0, %arg1 : tensor<?x?x?xf32>, tensor<?x?x?xf32>) + outs(%arg2 : tensor<?x?x?xf32>) { + ^bb0(%arg3: f32, %arg4: f32, %arg5: f32): // no predecessors + %m = mulf %arg3, %arg4 : f32 + %a = addf %arg5, %m : f32 + linalg.yield %a : f32 + } -> tensor<?x?x?xf32> + return %0 : tensor<?x?x?xf32> +}
diff --git a/iree/compiler/Dialect/HAL/Conversion/HALToVM/ConvertCommandBufferOps.cpp b/iree/compiler/Dialect/HAL/Conversion/HALToVM/ConvertCommandBufferOps.cpp index 3ff4c80..0ac4bc9 100644 --- a/iree/compiler/Dialect/HAL/Conversion/HALToVM/ConvertCommandBufferOps.cpp +++ b/iree/compiler/Dialect/HAL/Conversion/HALToVM/ConvertCommandBufferOps.cpp
@@ -74,6 +74,14 @@ patterns.insert<VMImportOpConversion<IREE::HAL::CommandBufferEndOp>>( context, importSymbols, typeConverter, "hal.command_buffer.end"); patterns + .insert<VMImportOpConversion<IREE::HAL::CommandBufferBeginDebugGroupOp>>( + context, importSymbols, typeConverter, + "hal.command_buffer.begin_debug_group"); + patterns + .insert<VMImportOpConversion<IREE::HAL::CommandBufferEndDebugGroupOp>>( + context, importSymbols, typeConverter, + "hal.command_buffer.end_debug_group"); + patterns .insert<VMImportOpConversion<IREE::HAL::CommandBufferExecutionBarrierOp>>( context, importSymbols, typeConverter, "hal.command_buffer.execution_barrier");
diff --git a/iree/compiler/Dialect/HAL/IR/HALOps.td b/iree/compiler/Dialect/HAL/IR/HALOps.td index 158182a..cc68781 100644 --- a/iree/compiler/Dialect/HAL/IR/HALOps.td +++ b/iree/compiler/Dialect/HAL/IR/HALOps.td
@@ -1125,6 +1125,43 @@ let hasCanonicalizer = 1; } +def HAL_CommandBufferBeginDebugGroupOp : HAL_Op<"command_buffer.begin_debug_group"> { + let summary = [{pushes a command buffer debug group label}]; + let description = [{ + Pushes a new debug group with the given label. + All commands between this and a mandatory matching call to + `hal.command_buffer.end_debug_group` will be grouped together with the + given label. + }]; + + let arguments = (ins + HAL_CommandBuffer:$command_buffer, + StrAttr:$label + ); + + let assemblyFormat = [{ + `<` $command_buffer `:` type($command_buffer) `>` + `label` `(` $label `)` + attr-dict-with-keyword + }]; +} + +def HAL_CommandBufferEndDebugGroupOp : HAL_Op<"command_buffer.end_debug_group"> { + let summary = [{pops a command buffer debug group label}]; + let description = [{ + Pops a debug group from the stack. + }]; + + let arguments = (ins + HAL_CommandBuffer:$command_buffer + ); + + let assemblyFormat = [{ + `<` $command_buffer `:` type($command_buffer) `>` + attr-dict-with-keyword + }]; +} + def HAL_CommandBufferExecutionBarrierOp : HAL_Op<"command_buffer.execution_barrier"> { let summary = [{command buffer execution barrier recording operation}]; let description = [{
diff --git a/iree/compiler/Dialect/HAL/Target/ROCM/ROCMTarget.cpp b/iree/compiler/Dialect/HAL/Target/ROCM/ROCMTarget.cpp index 9457a92..584f6d0 100644 --- a/iree/compiler/Dialect/HAL/Target/ROCM/ROCMTarget.cpp +++ b/iree/compiler/Dialect/HAL/Target/ROCM/ROCMTarget.cpp
@@ -39,8 +39,13 @@ llvm::cl::desc("Whether to try Linking to AMD Bitcodes"), llvm::cl::init(false)); + static llvm::cl::opt<std::string> clROCMBitcodeDir( + "iree-rocm-bc-dir", llvm::cl::desc("Directory of ROCM Bitcode"), + llvm::cl::init("/opt/rocm/amdgcn/bitcode")); + targetOptions.ROCMTargetChip = clROCMTargetChip; targetOptions.ROCMLinkBC = clROCMLinkBC; + targetOptions.ROCMBitcodeDir = clROCMBitcodeDir; return targetOptions; } @@ -150,7 +155,8 @@ // Link module to Device Library if (options_.ROCMLinkBC) - LinkROCDLIfNecessary(llvmModule.get(), options_.ROCMTargetChip); + LinkROCDLIfNecessary(llvmModule.get(), options_.ROCMTargetChip, + options_.ROCMBitcodeDir); // Serialize hsaco kernel into the binary that we will embed in the // final flatbuffer.
diff --git a/iree/compiler/Dialect/HAL/Target/ROCM/ROCMTarget.h b/iree/compiler/Dialect/HAL/Target/ROCM/ROCMTarget.h index 5a28c06..d4aa445 100644 --- a/iree/compiler/Dialect/HAL/Target/ROCM/ROCMTarget.h +++ b/iree/compiler/Dialect/HAL/Target/ROCM/ROCMTarget.h
@@ -20,6 +20,8 @@ std::string ROCMTargetChip; // Whether to try Linking to AMD Bitcodes bool ROCMLinkBC; + // AMD Bitcodes Directory + std::string ROCMBitcodeDir; }; ROCMTargetOptions getROCMTargetOptionsFromFlags(); @@ -29,7 +31,8 @@ std::function<ROCMTargetOptions()> queryOptions); // Links LLVM module to ROC Device Library Bit Code -void LinkROCDLIfNecessary(llvm::Module *module, std::string targetChip); +void LinkROCDLIfNecessary(llvm::Module *module, std::string targetChip, + std::string bitCodeDir); // Compiles ISAToHsaco Code std::string createHsaco(const std::string isa, StringRef name);
diff --git a/iree/compiler/Dialect/HAL/Target/ROCM/ROCMTargetUtils.cpp b/iree/compiler/Dialect/HAL/Target/ROCM/ROCMTargetUtils.cpp index d4b798b..37a605b 100644 --- a/iree/compiler/Dialect/HAL/Target/ROCM/ROCMTargetUtils.cpp +++ b/iree/compiler/Dialect/HAL/Target/ROCM/ROCMTargetUtils.cpp
@@ -90,7 +90,8 @@ return success(); } -static std::vector<std::string> GetROCDLPaths(std::string targetChip) { +static std::vector<std::string> GetROCDLPaths(std::string targetChip, + std::string bitCodeDir) { // AMDGPU bitcodes. int lenOfChipPrefix = 3; std::string chipId = targetChip.substr(lenOfChipPrefix); @@ -101,22 +102,22 @@ "oclc_daz_opt_off.bc", chip_isa_bc}); // Construct full path to ROCDL bitcode libraries. - std::string rocdl_dir_path = "/opt/rocm/amdgcn/bitcode"; std::vector<std::string> result; std::string app = "/"; for (auto &filename : *rocdl_filenames) { - result.push_back(rocdl_dir_path + app + filename); + result.push_back(bitCodeDir + app + filename); } return result; } // Links ROCm-Device-Libs into the given module if the module needs it. -void LinkROCDLIfNecessary(llvm::Module *module, std::string targetChip) { +void LinkROCDLIfNecessary(llvm::Module *module, std::string targetChip, + std::string bitCodeDir) { if (!HAL::CouldNeedDeviceBitcode(*module)) { return; } - if (!succeeded( - HAL::LinkWithBitcodeVector(module, GetROCDLPaths(targetChip)))) { + if (!succeeded(HAL::LinkWithBitcodeVector( + module, GetROCDLPaths(targetChip, bitCodeDir)))) { llvm::WithColor::error(llvm::errs()) << "Fail to Link ROCDL.\n"; }; }
diff --git a/iree/compiler/Dialect/HAL/hal.imports.mlir b/iree/compiler/Dialect/HAL/hal.imports.mlir index 2470eab..c51db3f 100644 --- a/iree/compiler/Dialect/HAL/hal.imports.mlir +++ b/iree/compiler/Dialect/HAL/hal.imports.mlir
@@ -141,6 +141,17 @@ %command_buffer : !vm.ref<!hal.command_buffer> ) +// Pushes a new debug group with the given |label|. +vm.import @command_buffer.begin_debug_group( + %command_buffer : !vm.ref<!hal.command_buffer>, + %label : !vm.buffer +) + +// Pops a debug group from the stack. +vm.import @command_buffer.end_debug_group( + %command_buffer : !vm.ref<!hal.command_buffer> +) + // Defines an execution dependency between all commands recorded before the // barrier and all commands recorded after the barrier. Only the stages provided // will be affected.
diff --git a/iree/compiler/Dialect/LinalgExt/IR/LinalgExtOps.cpp b/iree/compiler/Dialect/LinalgExt/IR/LinalgExtOps.cpp index eb0c002..2a9256c 100644 --- a/iree/compiler/Dialect/LinalgExt/IR/LinalgExtOps.cpp +++ b/iree/compiler/Dialect/LinalgExt/IR/LinalgExtOps.cpp
@@ -254,7 +254,8 @@ Operation *ScatterOp::getTiledImplementation( OpBuilder &builder, ValueRange outputs, ArrayRef<OpFoldResult> offsets, ArrayRef<OpFoldResult> sizes, - SmallVectorImpl<SmallVector<OpFoldResult, 4>> &resultOffsets) { + SmallVectorImpl<SmallVector<OpFoldResult, 4>> &resultOffsets, + SmallVectorImpl<SmallVector<OpFoldResult, 4>> &resultSizes) { assert(outputs.size() == 1 && offsets.size() == 1 && sizes.size() == 1); Location loc = getLoc(); auto zeroAttr = builder.getI64IntegerAttr(0); @@ -289,7 +290,12 @@ assert(tiledIndices && "failed to get slice of indices"); resultOffsets.resize(1); - resultOffsets[0].resize(getUpdateType().getRank(), zeroAttr); + resultOffsets[0].resize(updateRank, zeroAttr); + resultSizes.resize(1); + resultSizes[0].resize(updateRank); + for (auto dim : llvm::seq<int64_t>(0, updateRank)) { + resultSizes[0][dim] = getDim(builder, loc, original(), dim); + } SmallVector<Type> resultTypes; if (getNumResults()) { resultTypes.push_back(getResultTypes()[0]); @@ -400,7 +406,8 @@ Operation *SortOp::getTiledImplementation( OpBuilder &builder, ValueRange outputs, ArrayRef<OpFoldResult> offsets, ArrayRef<OpFoldResult> sizes, - SmallVectorImpl<SmallVector<OpFoldResult, 4>> &resultOffsets) { + SmallVectorImpl<SmallVector<OpFoldResult, 4>> &resultOffsets, + SmallVectorImpl<SmallVector<OpFoldResult, 4>> &resultSizes) { assert(outputs.size() == this->outputs().size()); int64_t rank = getOperandRank(); assert(offsets.size() == static_cast<size_t>(rank) && @@ -410,11 +417,13 @@ Location loc = getLoc(); SmallVector<Value> tiledOperands(outputs.size()); resultOffsets.resize(outputs.size()); + resultSizes.resize(outputs.size()); for (auto en : llvm::enumerate(outputs)) { tiledOperands[en.index()] = getSlice(builder, getLoc(), en.value(), offsets, sizes, strides); assert(tiledOperands[en.index()] && "failed to get slice of operand"); resultOffsets[en.index()].assign(offsets.begin(), offsets.end()); + resultSizes[en.index()].assign(sizes.begin(), sizes.end()); } SmallVector<Type, 4> resultTypes; if (getNumResults()) {
diff --git a/iree/compiler/Dialect/LinalgExt/IR/TiledOpInterface.td b/iree/compiler/Dialect/LinalgExt/IR/TiledOpInterface.td index 697acae..77f6c5f 100644 --- a/iree/compiler/Dialect/LinalgExt/IR/TiledOpInterface.td +++ b/iree/compiler/Dialect/LinalgExt/IR/TiledOpInterface.td
@@ -49,7 +49,8 @@ "ValueRange ":$outputs, "ArrayRef<OpFoldResult> ":$offsets, "ArrayRef<OpFoldResult> ":$sizes, - "SmallVectorImpl<SmallVector<OpFoldResult, 4>> &":$resultOffsets), + "SmallVectorImpl<SmallVector<OpFoldResult, 4>> &":$resultOffsets, + "SmallVectorImpl<SmallVector<OpFoldResult, 4>> &":$resultSizes), /*methodBody=*/"", /*defaultImplementation=*/[{ return nullptr;
diff --git a/iree/compiler/Dialect/LinalgExt/Transforms/Passes.h b/iree/compiler/Dialect/LinalgExt/Transforms/Passes.h index d84b97e..cac06fd 100644 --- a/iree/compiler/Dialect/LinalgExt/Transforms/Passes.h +++ b/iree/compiler/Dialect/LinalgExt/Transforms/Passes.h
@@ -13,7 +13,7 @@ namespace iree_compiler { namespace linalg_ext { -std::unique_ptr<OperationPass<FuncOp>> createLinalgExtTilingPass(); +std::unique_ptr<OperationPass<FuncOp>> createTiledOpInterfaceTilingPass(); std::unique_ptr<OperationPass<FuncOp>> createLinalgExtToLoopsPass();
diff --git a/iree/compiler/Dialect/LinalgExt/Transforms/Passes.td b/iree/compiler/Dialect/LinalgExt/Transforms/Passes.td index e221e33..9497cdb 100644 --- a/iree/compiler/Dialect/LinalgExt/Transforms/Passes.td +++ b/iree/compiler/Dialect/LinalgExt/Transforms/Passes.td
@@ -15,10 +15,10 @@ let constructor = "mlir::iree_compiler::linalg_ext::createLinalgExtToLoopsPass()"; } -def LinalgExtTiling : +def TiledOpInterfaceTiling : Pass<"iree-linalg-ext-tile", "FuncOp"> { - let summary = "Test pass for tiling LinalgExt ops using TiledOpInterface"; - let constructor = "mlir::iree_compiler::linalg_ext::createLinalgExtTilingPass()"; + let summary = "Test pass for tiling using TiledOpInterface"; + let constructor = "mlir::iree_compiler::linalg_ext::createTiledOpInterfaceTilingPass()"; } #endif // IREE_DIALECT_LINALGEXT_PASSES
diff --git a/iree/compiler/Dialect/LinalgExt/Transforms/Tiling.cpp b/iree/compiler/Dialect/LinalgExt/Transforms/Tiling.cpp index 7b4be2c..f686600 100644 --- a/iree/compiler/Dialect/LinalgExt/Transforms/Tiling.cpp +++ b/iree/compiler/Dialect/LinalgExt/Transforms/Tiling.cpp
@@ -82,13 +82,26 @@ return valueOrAttr.get<Value>(); } +/// Returns the constant value in `valueOrAttr` if it is not a dynamic `Value`. +static Optional<int64_t> getConstantValue(OpFoldResult valueOrAttr) { + if (auto attr = valueOrAttr.dyn_cast<Attribute>()) { + return attr.cast<IntegerAttr>().getInt(); + } + return {}; +} + +/// Checks if `valueOrAttr` represents a constant value `val`. +static bool isValue(OpFoldResult valueOrAttr, int64_t val) { + auto attr = valueOrAttr.dyn_cast<Attribute>(); + return attr && attr.cast<IntegerAttr>().getValue() == val; +} + /// Returns true if loop is untiled. Only checks if the value is statically /// zero. It is assumed that a `Value` defined by a constant op is already /// converted to an `IntegerAttr` of that value. So here just return true if /// this is an attribute with a zero value. static bool isUntiledLoop(OpFoldResult valueOrAttr) { - auto attr = valueOrAttr.dyn_cast<Attribute>(); - return attr && attr.cast<IntegerAttr>().getValue() == 0; + return isValue(valueOrAttr, 0); } /// Generates the tiled loops and the body by invoking the interface methods of @@ -108,22 +121,23 @@ /// distributed loops. It is a stack, and once an entry at the top of the /// stack is used for distribution it is popped before processing the inner /// loops. -static FailureOr<TiledOp> tileLinalgExtOpImpl( - OpBuilder &builder, TiledOpInterface op, ValueRange outputs, +static FailureOr<TiledOp> tileInterfaceOpImpl( + OpBuilder &builder, TiledOpInterface tilableOp, ValueRange outputs, MutableArrayRef<OpFoldResult> tileSizes, ArrayRef<StringRef> iteratorTypes, ArrayRef<Range> loopBounds, unsigned loopDepth, SmallVectorImpl<OpFoldResult> &offsets, ArrayRef<linalg::ProcInfo> distributionInfo) { - Location loc = op.getLoc(); + Location loc = tilableOp.getLoc(); // If this is the innermost loop, then generated the tiled implementation of // the op by invoking the TiledOpInterface methods. if (loopDepth == tileSizes.size()) { SmallVector<SmallVector<OpFoldResult, 4>> resultOffsets; - Operation *tiledOp = op.getTiledImplementation(builder, outputs, offsets, - tileSizes, resultOffsets); + SmallVector<SmallVector<OpFoldResult, 4>> resultSizes; + Operation *tiledOp = tilableOp.getTiledImplementation( + builder, outputs, offsets, tileSizes, resultOffsets, resultSizes); if (!tiledOp) { return static_cast<LogicalResult>( - op.emitOpError("failed to get tiled implementation")); + tilableOp.emitOpError("failed to get tiled implementation")); } assert(tiledOp->getNumResults() == 0 || (resultOffsets.size() == tiledOp->getNumResults())); @@ -134,16 +148,13 @@ // to be inserted into the `initValues` and returned. if (tiledOp->getNumResults()) { SmallVector<Value> results; + auto oneAttr = builder.getI64IntegerAttr(1); results.reserve(tiledOp->getNumResults()); for (auto en : llvm::enumerate(tiledOp->getResults())) { Value result = en.value(); ArrayRef<OpFoldResult> offsets(resultOffsets[en.index()]); - auto resultType = result.getType().cast<ShapedType>(); - auto oneAttr = builder.getI64IntegerAttr(1); - SmallVector<OpFoldResult> strides(resultType.getRank(), oneAttr); - auto sizes = llvm::to_vector<4>(llvm::map_range( - llvm::seq<int64_t>(0, resultType.getRank()), - [&](int64_t dim) { return getDim(builder, loc, result, dim); })); + ArrayRef<OpFoldResult> sizes(resultSizes[en.index()]); + SmallVector<OpFoldResult> strides(offsets.size(), oneAttr); Value insert = builder.create<tensor::InsertSliceOp>( loc, result, outputs[en.index()], offsets, sizes, strides); results.push_back(insert); @@ -160,9 +171,9 @@ assert(matchPattern(loopBounds[loopDepth].offset, m_Zero()) && "expected loop bounds to have lower bound of zero"); tileSizes[loopDepth] = getOpFoldResult(loopBounds[loopDepth].size); - return tileLinalgExtOpImpl(builder, op, outputs, tileSizes, iteratorTypes, - loopBounds, loopDepth + 1, offsets, - distributionInfo); + return tileInterfaceOpImpl(builder, tilableOp, outputs, tileSizes, + iteratorTypes, loopBounds, loopDepth + 1, + offsets, distributionInfo); } // Generate an scf.for for the current loop depth. @@ -170,7 +181,7 @@ Value ub = loopBounds[loopDepth].size; if (!matchPattern(loopBounds[loopDepth].stride, m_One())) { return static_cast<LogicalResult>( - op.emitOpError("expected stride to be 1")); + tilableOp.emitOpError("expected stride to be 1")); } Value step = getValue(builder, loc, tileSizes[loopDepth]); @@ -183,7 +194,7 @@ distributionInfo = distributionInfo.drop_front(); } FailureOr<TiledOp> innerReturnValue; - bool isBufferTiling = op->getNumResults() == 0; + bool isBufferTiling = tilableOp->getNumResults() == 0; ValueRange initValues(isBufferTiling ? ValueRange{} : outputs); auto forOp = builder.create<scf::ForOp>( loc, lb, ub, step, initValues, @@ -200,9 +211,10 @@ ValueRange{iv, getValue(builder, loc, tileSizes[loopDepth]), ub}); tileSizes[loopDepth] = getOpFoldResult(inBoundsTileSize); // Recursively proceed to generate the tiled loop for the next level. - innerReturnValue = tileLinalgExtOpImpl( - b, op, (isBufferTiling ? outputs : args), tileSizes, iteratorTypes, - loopBounds, loopDepth + 1, offsets, distributionInfo); + innerReturnValue = + tileInterfaceOpImpl(b, tilableOp, (isBufferTiling ? outputs : args), + tileSizes, iteratorTypes, loopBounds, + loopDepth + 1, offsets, distributionInfo); if (failed(innerReturnValue)) return; b.create<scf::YieldOp>(loc, innerReturnValue->results); }); @@ -215,14 +227,14 @@ return innerReturnValue; } -FailureOr<TiledOp> tileLinalgExtOp(OpBuilder &b, LinalgExtOp op, +FailureOr<TiledOp> tileInterfaceOp(OpBuilder &b, Operation *op, ValueRange dest, const linalg::LinalgTilingOptions &options) { - TiledOpInterface tilableOp = dyn_cast<TiledOpInterface>(op.getOperation()); + TiledOpInterface tilableOp = dyn_cast<TiledOpInterface>(op); if (!tilableOp) return TiledOp{}; SmallVector<StringRef> iteratorTypes = tilableOp.getLoopIteratorTypes(); SmallVector<Value, 4> tileSizesVals = - options.tileSizeComputationFunction(b, tilableOp.getOperation()); + options.tileSizeComputationFunction(b, op); auto zeroAttr = b.getI64IntegerAttr(0); // The actual tile sizes used converts `Value` defined as constant 0, to a @@ -233,14 +245,14 @@ for (auto en : llvm::enumerate(iteratorTypes)) { if (en.value() == getParallelIteratorTypeName()) continue; if (!isUntiledLoop(tileSizes[en.index()])) { - return static_cast<LogicalResult>(op.emitOpError( + return static_cast<LogicalResult>(tilableOp.emitOpError( "unimplemented tiling of non-parallel loop iterator type")); } } // Trivial early exit case of tile sizes being zero for all parallel loops. if (llvm::all_of(tileSizes, isUntiledLoop)) { - return TiledOp{op.getOperation(), {}, {}}; + return TiledOp{op, {}, {}}; } SmallVector<Range> loopBounds = tilableOp.getLoopBounds(b); @@ -257,33 +269,130 @@ if (iteratorTypes[i] != getParallelIteratorTypeName()) continue; distributedLoopRange.push_back(loopBounds[i]); } - distributionInfo = - options.distribution->procInfo(b, op.getLoc(), distributedLoopRange); + distributionInfo = options.distribution->procInfo(b, tilableOp.getLoc(), + distributedLoopRange); } SmallVector<OpFoldResult> offsets; - return tileLinalgExtOpImpl(b, tilableOp, op.outputs(), tileSizes, - iteratorTypes, loopBounds, 0, offsets, - distributionInfo); + return tileInterfaceOpImpl(b, tilableOp, dest, tileSizes, iteratorTypes, + loopBounds, 0, offsets, distributionInfo); } //===----------------------------------------------------------------------===// +// Interface implementations for external operations. +//===----------------------------------------------------------------------===// + +namespace { +struct InsertSliceTiledOpInterface + : public TiledOpInterface::ExternalModel<InsertSliceTiledOpInterface, + tensor::InsertSliceOp> { + SmallVector<StringRef> getLoopIteratorTypes(Operation *op) const { + auto insertSliceOp = cast<tensor::InsertSliceOp>(op); + return SmallVector<StringRef>(insertSliceOp.getSourceType().getRank(), + getParallelIteratorTypeName()); + } + + SmallVector<Range> getLoopBounds(Operation *op, OpBuilder &b) const { + auto insertSliceOp = cast<tensor::InsertSliceOp>(op); + Value source = insertSliceOp.source(); + RankedTensorType sourceType = insertSliceOp.getSourceType(); + Location loc = op->getLoc(); + Value zero = b.create<ConstantIndexOp>(loc, 0); + Value one = b.create<ConstantIndexOp>(loc, 1); + SmallVector<Range> loopBounds(sourceType.getRank(), + Range{zero, nullptr, one}); + for (auto dim : + llvm::seq<int64_t>(0, insertSliceOp.getSourceType().getRank())) { + loopBounds[dim].size = b.create<tensor::DimOp>(loc, source, dim); + } + return loopBounds; + } + + Operation *getTiledImplementation( + Operation *op, OpBuilder &b, ValueRange outputs, + ArrayRef<OpFoldResult> offsets, ArrayRef<OpFoldResult> sizes, + SmallVectorImpl<SmallVector<OpFoldResult, 4>> &resultOffsets, + SmallVectorImpl<SmallVector<OpFoldResult, 4>> &resultSizes) const { + // Compute a subtensor of the source based on the offsets. + auto insertOp = cast<tensor::InsertSliceOp>(op); + auto opOffsets = insertOp.getMixedOffsets(); + auto opSizes = insertOp.getMixedSizes(); + auto opStrides = insertOp.getMixedStrides(); + if (!llvm::all_of(opStrides, [&](OpFoldResult valueOrAttr) { + return isValue(valueOrAttr, 1); + })) { + op->emitOpError("unable to tile operation with non-unit stride"); + return nullptr; + } + // The operation returned is just a tensor.extract_slice of the source with + // the given offsets, sizes and strides. Setting the correct result offset + // will make the sure the tiling algorithm will insert this slice into the + // correct place in the destination. + // The result offset is just the offset passed in plus the offset specified + // in the op (since all strides are checked to be 1). + unsigned offsetIndex = 0; + ArrayRef<int64_t> sourceShape = insertOp.getSourceType().getShape(); + int64_t destRank = insertOp.getType().getRank(); + resultOffsets.resize(1); + resultOffsets[0].resize(destRank); + resultSizes.resize(1); + resultSizes[0].resize(destRank); + Location loc = insertOp.getLoc(); + auto zeroAttr = b.getI64IntegerAttr(0); + auto oneAttr = b.getI64IntegerAttr(1); + SmallVector<OpFoldResult> strides(offsets.size(), oneAttr); + for (auto opOffset : llvm::enumerate(opOffsets)) { + // Check for rank-reducing by checking that + // 1) The corresponding opSize value is 1 + // 2) The current rank of the source is not 1. + // Then the opOffset is for the rank-reduced dimension. Skip. + unsigned opOffsetIndex = opOffset.index(); + if (isValue(opSizes[opOffsetIndex], 1) && sourceShape[offsetIndex] != 1) { + resultOffsets[0][opOffsetIndex] = zeroAttr; + resultSizes[0][opOffsetIndex] = oneAttr; + continue; + } + OpFoldResult opOffsetVal = opOffset.value(); + OpFoldResult offset = offsets[offsetIndex]; + if (opOffsetVal.is<Attribute>() && offset.is<Attribute>()) { + resultOffsets[0][opOffsetIndex] = b.getI64IntegerAttr( + *getConstantValue(opOffsetVal) + *getConstantValue(offset)); + } else { + AffineMap map = AffineMap::get( + 1, 1, {b.getAffineDimExpr(0) + b.getAffineSymbolExpr(0)}); + resultOffsets[0][opOffsetIndex] = + b.create<AffineApplyOp>(loc, map, + ValueRange{getValue(b, loc, offset), + getValue(b, loc, opOffsetVal)}) + .getResult(); + } + resultSizes[0][opOffsetIndex] = sizes[offsetIndex]; + offsetIndex++; + } + return b.create<tensor::ExtractSliceOp>(loc, insertOp.source(), offsets, + sizes, strides); + } +}; +} // namespace + +//===----------------------------------------------------------------------===// // Patterns for tiling LinalgExtOps. //===----------------------------------------------------------------------===// namespace { -/// Base pattern for tiling LinalgExtOps. -struct LinalgExtBaseTilingPattern : public RewritePattern { - LinalgExtBaseTilingPattern(StringRef opName, MLIRContext *context, - linalg::LinalgTilingOptions options, - linalg::LinalgTransformationFilter filter = - linalg::LinalgTransformationFilter(), - PatternBenefit benefit = 1) +/// Base pattern for tiling TiledOpInterfaceOps. +struct TiledOpInterfaceBaseTilingPattern : public RewritePattern { + TiledOpInterfaceBaseTilingPattern(StringRef opName, MLIRContext *context, + linalg::LinalgTilingOptions options, + linalg::LinalgTransformationFilter filter = + linalg::LinalgTransformationFilter(), + PatternBenefit benefit = 1) : RewritePattern(opName, benefit, context), filter(filter), options(options) {} - LogicalResult matchAndRewriteBase(Operation *op, PatternRewriter &rewriter, + LogicalResult matchAndRewriteBase(Operation *op, ValueRange dest, + PatternRewriter &rewriter, TiledOp &result) const; private: @@ -294,21 +403,54 @@ }; template <typename OpTy> -struct LinalgExtTilingPattern : public LinalgExtBaseTilingPattern { +struct LinalgExtTilingPattern : public TiledOpInterfaceBaseTilingPattern { LinalgExtTilingPattern(MLIRContext *context, linalg::LinalgTilingOptions options, linalg::LinalgTransformationFilter filter = linalg::LinalgTransformationFilter(), PatternBenefit benefit = 1) - : LinalgExtBaseTilingPattern(OpTy::getOperationName(), context, options, - filter, benefit) {} + : TiledOpInterfaceBaseTilingPattern(OpTy::getOperationName(), context, + options, filter, benefit) {} LogicalResult matchAndRewrite(Operation *op, PatternRewriter &rewriter) const override { + auto linalgExtOp = cast<LinalgExtOp>(op); TiledOp tiledOp; // Check for failure. - if (failed(LinalgExtBaseTilingPattern::matchAndRewriteBase(op, rewriter, - tiledOp))) { + if (failed(TiledOpInterfaceBaseTilingPattern::matchAndRewriteBase( + op, linalgExtOp.outputs(), rewriter, tiledOp))) { + return failure(); + } + // Check for do-nothing case. + if (!tiledOp.op) return failure(); + if (tiledOp.op != op) { + if (tiledOp.results.empty()) { + rewriter.eraseOp(op); + } else { + rewriter.replaceOp(op, tiledOp.results); + } + } + return success(); + } +}; + +struct InsertSliceTilingPattern : public TiledOpInterfaceBaseTilingPattern { + InsertSliceTilingPattern(MLIRContext *context, + linalg::LinalgTilingOptions options, + linalg::LinalgTransformationFilter filter = + linalg::LinalgTransformationFilter(), + PatternBenefit benefit = 1) + : TiledOpInterfaceBaseTilingPattern( + tensor::InsertSliceOp::getOperationName(), context, options, filter, + benefit) {} + + LogicalResult matchAndRewrite(Operation *op, + PatternRewriter &rewriter) const override { + tensor::InsertSliceOp insertSliceOp = cast<tensor::InsertSliceOp>(op); + TiledOp tiledOp; + // Check for failure. + if (failed(TiledOpInterfaceBaseTilingPattern::matchAndRewriteBase( + insertSliceOp, insertSliceOp.dest(), rewriter, tiledOp))) { return failure(); } // Check for do-nothing case. @@ -325,16 +467,15 @@ }; } // namespace -LogicalResult LinalgExtBaseTilingPattern::matchAndRewriteBase( - Operation *op, PatternRewriter &rewriter, TiledOp &result) const { - auto linalgExtOp = dyn_cast<LinalgExtOp>(op); - if (!linalgExtOp) return failure(); +LogicalResult TiledOpInterfaceBaseTilingPattern::matchAndRewriteBase( + Operation *op, ValueRange dest, PatternRewriter &rewriter, + TiledOp &result) const { if (failed(filter.checkAndNotify(rewriter, op))) return failure(); if (failed(verifySupportedTilingOptions(rewriter, op, options))) { return failure(); } - FailureOr<TiledOp> res = tileLinalgExtOp(rewriter, linalgExtOp, options); + FailureOr<TiledOp> res = tileInterfaceOp(rewriter, op, dest, options); if (failed(res)) return res; result = *res; if (result.op) { @@ -348,7 +489,8 @@ //===----------------------------------------------------------------------===// namespace { -struct LinalgExtTilingPass : public LinalgExtTilingBase<LinalgExtTilingPass> { +struct TiledOpInterfaceTilingPass + : public TiledOpInterfaceTilingBase<TiledOpInterfaceTilingPass> { void getDependentDialects(DialectRegistry ®istry) const override { registry .insert<AffineDialect, IREE::Flow::FlowDialect, linalg::LinalgDialect, @@ -356,6 +498,7 @@ tensor::TensorDialect, scf::SCFDialect>(); } + LogicalResult initialize(MLIRContext *context) override; void runOnOperation() override; }; } // namespace @@ -365,9 +508,17 @@ return b.template create<OpTy>(b.getInsertionPoint()->getLoc(), dim); } -void LinalgExtTilingPass::runOnOperation() { +LogicalResult TiledOpInterfaceTilingPass::initialize(MLIRContext *context) { + tensor::InsertSliceOp::attachInterface<InsertSliceTiledOpInterface>(*context); + return success(); +} + +void TiledOpInterfaceTilingPass::runOnOperation() { FuncOp funcOp = getOperation(); MLIRContext *context = funcOp.getContext(); + + tensor::InsertSliceOp::attachInterface<InsertSliceTiledOpInterface>(*context); + RewritePatternSet patterns(context); patterns.add<LinalgExtTilingPattern<ScatterOp>>( context, linalg::LinalgTilingOptions().setTileSizes({10, 20}), @@ -379,6 +530,7 @@ linalg::LinalgTransformationFilter( Identifier::get("no_tiling_input", context), Identifier::get("no_tiling_output", context))); + patterns.add<LinalgExtTilingPattern<SortOp>>( context, linalg::LinalgTilingOptions().setTileSizes({0, 20}), linalg::LinalgTransformationFilter( @@ -419,13 +571,19 @@ Identifier::get("distribute_input", context), Identifier::get("distribute_output", context))); + patterns.add<InsertSliceTilingPattern>( + context, linalg::LinalgTilingOptions().setTileSizes({10, 20}), + linalg::LinalgTransformationFilter( + Identifier::get("tiling_input", context), + Identifier::get("tiling_output", context))); + if (failed(applyPatternsAndFoldGreedily(funcOp, std::move(patterns)))) { return signalPassFailure(); } } -std::unique_ptr<OperationPass<FuncOp>> createLinalgExtTilingPass() { - return std::make_unique<LinalgExtTilingPass>(); +std::unique_ptr<OperationPass<FuncOp>> createTiledOpInterfaceTilingPass() { + return std::make_unique<TiledOpInterfaceTilingPass>(); } } // namespace linalg_ext
diff --git a/iree/compiler/Dialect/LinalgExt/Transforms/Transforms.h b/iree/compiler/Dialect/LinalgExt/Transforms/Transforms.h index 08e1a81..e251cab 100644 --- a/iree/compiler/Dialect/LinalgExt/Transforms/Transforms.h +++ b/iree/compiler/Dialect/LinalgExt/Transforms/Transforms.h
@@ -27,7 +27,7 @@ /// Main entry point for tiling LinalgExtOps using TiledOpInterface. If the /// `op` does not implement the `TiledOpInterface` returns a `TiledOp{}` value. -FailureOr<TiledOp> tileLinalgExtOp(OpBuilder &b, LinalgExtOp op, +FailureOr<TiledOp> tileLinalgExtOp(OpBuilder &b, Operation *op, ValueRange dest, const linalg::LinalgTilingOptions &options); } // namespace linalg_ext
diff --git a/iree/compiler/Dialect/LinalgExt/Transforms/test/tiling.mlir b/iree/compiler/Dialect/LinalgExt/Transforms/test/tiling.mlir index b870530..5818c01 100644 --- a/iree/compiler/Dialect/LinalgExt/Transforms/test/tiling.mlir +++ b/iree/compiler/Dialect/LinalgExt/Transforms/test/tiling.mlir
@@ -30,12 +30,12 @@ // CHECK-SAME: [%[[USED_TILESIZE]], %[[D1]]] // CHECK: %[[INDEX_SLICE:.+]] = tensor.extract_slice %[[INDICES]][%[[IV]], 0] // CHECK-SAME: [%[[USED_TILESIZE]], 1] +// CHECK-DAG: %[[SLICE_D0:.+]] = tensor.dim %[[ORIGINAL]], %[[C0]] +// CHECK-DAG: %[[SLICE_D1:.+]] = tensor.dim %[[ORIGINAL]], %[[C1]] // CHECK: %[[SCATTER_TILE:.+]] = linalg_ext.scatter // CHECK-SAME: __internal_linalg_transform__ = "tiling_output" // CHECK-SAME: ins(%[[UPDATE_SLICE]], %[[INDEX_SLICE]] // CHECK-SAME: outs(%[[INIT]] -// CHECK-DAG: %[[SLICE_D0:.+]] = tensor.dim %[[SCATTER_TILE]], %[[C0]] -// CHECK-DAG: %[[SLICE_D1:.+]] = tensor.dim %[[SCATTER_TILE]], %[[C1]] // CHECK: %[[YIELD:.+]] = tensor.insert_slice %[[SCATTER_TILE]] into %[[INIT]][0, 0] // CHECK-SAME: [%[[SLICE_D0]], %[[SLICE_D1]]] // CHECK: scf.yield %[[YIELD]] @@ -114,12 +114,12 @@ // CHECK-SAME: [%[[USED_TILESIZE]], %[[D1]]] // CHECK: %[[INDEX_SLICE:.+]] = tensor.extract_slice %[[INDICES]][%[[IV]], 0] // CHECK-SAME: [%[[USED_TILESIZE]], 1] +// CHECK-DAG: %[[SLICE_D0:.+]] = tensor.dim %[[ORIGINAL]], %[[C0]] +// CHECK-DAG: %[[SLICE_D1:.+]] = tensor.dim %[[ORIGINAL]], %[[C1]] // CHECK: %[[SCATTER_TILE:.+]] = linalg_ext.scatter // CHECK-SAME: __internal_linalg_transform__ = "distribute_output" // CHECK-SAME: ins(%[[UPDATE_SLICE]], %[[INDEX_SLICE]] // CHECK-SAME: outs(%[[INIT]] -// CHECK-DAG: %[[SLICE_D0:.+]] = tensor.dim %[[SCATTER_TILE]], %[[C0]] -// CHECK-DAG: %[[SLICE_D1:.+]] = tensor.dim %[[SCATTER_TILE]], %[[C1]] // CHECK: %[[YIELD:.+]] = tensor.insert_slice %[[SCATTER_TILE]] into %[[INIT]][0, 0] // CHECK-SAME: [%[[SLICE_D0]], %[[SLICE_D1]]] // CHECK: scf.yield %[[YIELD]] @@ -197,10 +197,8 @@ // CHECK: %[[SORT_TILE:.+]] = linalg_ext.sort // CHECK-SAME: __internal_linalg_transform__ = "inner_reduce_output" // CHECK-SAME: outs(%[[OPERAND_SLICE]] -// CHECK-DAG: %[[SLICE_D0:.+]] = tensor.dim %[[SORT_TILE]], %[[C0]] -// CHECK-DAG: %[[SLICE_D1:.+]] = tensor.dim %[[SORT_TILE]], %[[C1]] // CHECK: %[[YIELD:.+]] = tensor.insert_slice %[[SORT_TILE]] into %[[INIT]][%[[IV]], 0] -// CHECK-SAME: [%[[SLICE_D0]], %[[SLICE_D1]]] +// CHECK-SAME: [%[[USED_TILESIZE]], %[[D1]]] // CHECK: scf.yield %[[YIELD]] // CHECK: return %[[RESULT]] @@ -232,10 +230,8 @@ // CHECK: %[[SORT_TILE:.+]] = linalg_ext.sort // CHECK-SAME: __internal_linalg_transform__ = "outer_reduce_output" // CHECK-SAME: outs(%[[OPERAND_SLICE]] -// CHECK-DAG: %[[SLICE_D0:.+]] = tensor.dim %[[SORT_TILE]], %[[C0]] -// CHECK-DAG: %[[SLICE_D1:.+]] = tensor.dim %[[SORT_TILE]], %[[C1]] // CHECK: %[[YIELD:.+]] = tensor.insert_slice %[[SORT_TILE]] into %[[INIT]][0, %[[IV]]] -// CHECK-SAME: [%[[SLICE_D0]], %[[SLICE_D1]]] +// CHECK-SAME: [%[[D0]], %[[USED_TILESIZE]]] // CHECK: scf.yield %[[YIELD]] // CHECK: return %[[RESULT]] @@ -272,14 +268,10 @@ // CHECK: %[[SORT_TILE:.+]]:2 = linalg_ext.sort // CHECK-SAME: __internal_linalg_transform__ = "inner_reduce_output" // CHECK-SAME: outs(%[[OPERAND1_SLICE]], %[[OPERAND2_SLICE]] -// CHECK-DAG: %[[SLICE_D0:.+]] = tensor.dim %[[SORT_TILE]]#0, %[[C0]] -// CHECK-DAG: %[[SLICE_D1:.+]] = tensor.dim %[[SORT_TILE]]#0, %[[C1]] // CHECK: %[[YIELD1:.+]] = tensor.insert_slice %[[SORT_TILE]]#0 into %[[INIT1]][%[[IV]], 0] -// CHECK-SAME: [%[[SLICE_D0]], %[[SLICE_D1]]] -// CHECK-DAG: %[[SLICE_D0_0:.+]] = tensor.dim %[[SORT_TILE]]#1, %[[C0]] -// CHECK-DAG: %[[SLICE_D1_0:.+]] = tensor.dim %[[SORT_TILE]]#1, %[[C1]] +// CHECK-SAME: [%[[USED_TILESIZE]], %[[D1]]] // CHECK: %[[YIELD2:.+]] = tensor.insert_slice %[[SORT_TILE]]#1 into %[[INIT2]][%[[IV]], 0] -// CHECK-SAME: [%[[SLICE_D0_0]], %[[SLICE_D1_0]]] +// CHECK-SAME: [%[[USED_TILESIZE]], %[[D1]]] // CHECK: scf.yield %[[YIELD1]], %[[YIELD2]] // CHECK: return %[[RESULT]]#0, %[[RESULT]]#1 @@ -420,3 +412,65 @@ // CHECK: linalg_ext.sort // CHECK-SAME: __internal_linalg_transform__ = "distribute_output" // CHECK-SAME: outs(%[[OPERAND1_SLICE]], %[[OPERAND2_SLICE]] + +// ----- + +func @slice_insert(%source :tensor<?x?xf32>, %dest: tensor<?x?xf32>, + %idx0 : index, %idx1 : index) -> tensor<?x?xf32> { + %c0 = constant 0 : index + %c1 = constant 1 : index + %0 = tensor.dim %source, %c0 : tensor<?x?xf32> + %1 = tensor.dim %source, %c1 : tensor<?x?xf32> + %2 = tensor.insert_slice %source into %dest[%idx0, %idx1] [%0, %1] [1, 1] + {__internal_linalg_transform__ = "tiling_input"} : tensor<?x?xf32> into tensor<?x?xf32> + return %2 : tensor<?x?xf32> +} +// CHECK-DAG: #[[MAP0:.+]] = affine_map<(d0)[s0, s1] -> (10, -d0 + s1)> +// CHECK-DAG: #[[MAP1:.+]] = affine_map<(d0)[s0, s1] -> (20, -d0 + s1)> +// CHECK-DAG: #[[MAP2:.+]] = affine_map<(d0)[s0] -> (d0 + s0)> +// CHECK: func @slice_insert( +// CHECK-SAME: %[[ARG0:[a-zA-Z0-9_]+]]: tensor<?x?xf32> +// CHECK-SAME: %[[ARG1:[a-zA-Z0-9_]+]]: tensor<?x?xf32> +// CHECK-SAME: %[[ARG2:[a-zA-Z0-9_]+]]: index +// CHECK-SAME: %[[ARG3:[a-zA-Z0-9_]+]]: index +// CHECK: %[[RESULT:.+]] = scf.for %[[IV0:[a-zA-Z0-9]+]] = +// CHECK: %[[YIELD1:.+]] = scf.for %[[IV1:[a-zA-Z0-9]+]] = +// CHECK: %[[OFFSET0:.+]] = affine.apply #[[MAP2]](%[[IV0]])[%[[ARG2]]] +// CHECK: %[[OFFSET1:.+]] = affine.apply #[[MAP2]](%[[IV1]])[%[[ARG3]]] +// CHECK: %[[SLICE:.+]] = tensor.extract_slice %[[ARG0]][%[[IV0]], %[[IV1]]] +// CHECK: %[[UPDATE:.+]] = tensor.insert_slice %[[SLICE]] +// CHECK-SAME: into %{{.+}}[%[[OFFSET0]], %[[OFFSET1]]] +// CHECK: scf.yield %[[UPDATE]] +// CHECK: scf.yield %[[YIELD1]] +// CHECK: return %[[RESULT]] + +// ----- + +func @slice_insert_rank_reduce(%source :tensor<?x?xf32>, %dest: tensor<?x?x?xf32>, + %idx0 : index, %idx1 : index) -> tensor<?x?x?xf32> { + %c0 = constant 0 : index + %c1 = constant 1 : index + %0 = tensor.dim %source, %c0 : tensor<?x?xf32> + %1 = tensor.dim %source, %c1 : tensor<?x?xf32> + %2 = tensor.insert_slice %source into %dest[%idx0, 0, %idx1] [%0, 1, %1] [1, 1, 1] + {__internal_linalg_transform__ = "tiling_input"} : tensor<?x?xf32> into tensor<?x?x?xf32> + return %2 : tensor<?x?x?xf32> +} +// CHECK-DAG: #[[MAP0:.+]] = affine_map<(d0)[s0, s1] -> (10, -d0 + s1)> +// CHECK-DAG: #[[MAP1:.+]] = affine_map<(d0)[s0, s1] -> (20, -d0 + s1)> +// CHECK-DAG: #[[MAP2:.+]] = affine_map<(d0)[s0] -> (d0 + s0)> +// CHECK: func @slice_insert_rank_reduce( +// CHECK-SAME: %[[ARG0:[a-zA-Z0-9_]+]]: tensor<?x?xf32> +// CHECK-SAME: %[[ARG1:[a-zA-Z0-9_]+]]: tensor<?x?x?xf32> +// CHECK-SAME: %[[ARG2:[a-zA-Z0-9_]+]]: index +// CHECK-SAME: %[[ARG3:[a-zA-Z0-9_]+]]: index +// CHECK: %[[RESULT:.+]] = scf.for %[[IV0:[a-zA-Z0-9]+]] = +// CHECK: %[[YIELD1:.+]] = scf.for %[[IV1:[a-zA-Z0-9]+]] = +// CHECK: %[[OFFSET0:.+]] = affine.apply #[[MAP2]](%[[IV0]])[%[[ARG2]]] +// CHECK: %[[OFFSET1:.+]] = affine.apply #[[MAP2]](%[[IV1]])[%[[ARG3]]] +// CHECK: %[[SLICE:.+]] = tensor.extract_slice %[[ARG0]][%[[IV0]], %[[IV1]]] +// CHECK: %[[UPDATE:.+]] = tensor.insert_slice %[[SLICE]] +// CHECK-SAME: into %{{.+}}[%[[OFFSET0]], 0, %[[OFFSET1]]] +// CHECK: scf.yield %[[UPDATE]] +// CHECK: scf.yield %[[YIELD1]] +// CHECK: return %[[RESULT]]
diff --git a/iree/hal/command_buffer.c b/iree/hal/command_buffer.c index 99b4758..c416684 100644 --- a/iree/hal/command_buffer.c +++ b/iree/hal/command_buffer.c
@@ -79,6 +79,22 @@ return status; } +IREE_API_EXPORT void iree_hal_command_buffer_begin_debug_group( + iree_hal_command_buffer_t* command_buffer, iree_string_view_t label, + iree_hal_label_color_t label_color, + const iree_hal_label_location_t* location) { + IREE_ASSERT_ARGUMENT(command_buffer); + _VTABLE_DISPATCH(command_buffer, begin_debug_group) + (command_buffer, label, label_color, location); +} + +IREE_API_EXPORT void iree_hal_command_buffer_end_debug_group( + iree_hal_command_buffer_t* command_buffer) { + IREE_ASSERT_ARGUMENT(command_buffer); + _VTABLE_DISPATCH(command_buffer, end_debug_group) + (command_buffer); +} + IREE_API_EXPORT iree_status_t iree_hal_command_buffer_execution_barrier( iree_hal_command_buffer_t* command_buffer, iree_hal_execution_stage_t source_stage_mask,
diff --git a/iree/hal/command_buffer.h b/iree/hal/command_buffer.h index 38d7242..82b3deb 100644 --- a/iree/hal/command_buffer.h +++ b/iree/hal/command_buffer.h
@@ -182,6 +182,26 @@ iree_device_size_t length; } iree_hal_buffer_barrier_t; +// An RGBA color. +typedef struct iree_hal_label_color_t { + uint8_t r; + uint8_t g; + uint8_t b; + uint8_t a; +} iree_hal_label_color_t; + +// A source location attached to debug labels. +typedef struct iree_hal_label_location_t { + iree_string_view_t file; + int line; +} iree_hal_label_location_t; + +// An unspecified color; debugging tools are to choose their own. +static inline iree_hal_label_color_t iree_hal_label_color_unspecified() { + iree_hal_label_color_t color = {0, 0, 0, 0}; + return color; +} + // TODO(benvanik): replace with tables for iree_string_builder_*. #define iree_hal_command_buffer_mode_string(...) "TODO" // {IREE_HAL_COMMAND_BUFFER_MODE_ONE_SHOT, "ONE_SHOT"}, @@ -262,6 +282,25 @@ IREE_API_EXPORT iree_status_t iree_hal_command_buffer_end(iree_hal_command_buffer_t* command_buffer); +// Pushes a new debug group with the given |label|. +// All commands between this and a mandatory matching call to +// iree_hal_command_buffer_end_debug_group will be grouped together with the +// given label. If a source location is available it can be provided via +// |location| to allow mapping back into the source program that issued the +// commands. +// +// An optional RGBA color to show in the debug UI may be provided via +// |label_color|; otherwise iree_hal_label_color_unspecified can be used to let +// the debug tool choose. +IREE_API_EXPORT void iree_hal_command_buffer_begin_debug_group( + iree_hal_command_buffer_t* command_buffer, iree_string_view_t label, + iree_hal_label_color_t label_color, + const iree_hal_label_location_t* location); + +// Pops a debug group from the stack. +IREE_API_EXPORT void iree_hal_command_buffer_end_debug_group( + iree_hal_command_buffer_t* command_buffer); + // Defines a memory dependency between commands recorded before and after the // barrier. One or more memory or buffer barriers can be specified to indicate // between which stages or buffers the dependencies exist. @@ -468,6 +507,13 @@ iree_status_t(IREE_API_PTR* begin)(iree_hal_command_buffer_t* command_buffer); iree_status_t(IREE_API_PTR* end)(iree_hal_command_buffer_t* command_buffer); + void(IREE_API_PTR* begin_debug_group)( + iree_hal_command_buffer_t* command_buffer, iree_string_view_t label, + iree_hal_label_color_t label_color, + const iree_hal_label_location_t* location); + void(IREE_API_PTR* end_debug_group)( + iree_hal_command_buffer_t* command_buffer); + iree_status_t(IREE_API_PTR* execution_barrier)( iree_hal_command_buffer_t* command_buffer, iree_hal_execution_stage_t source_stage_mask,
diff --git a/iree/hal/cuda/graph_command_buffer.c b/iree/hal/cuda/graph_command_buffer.c index 9b7fdc9..1b38e9d 100644 --- a/iree/hal/cuda/graph_command_buffer.c +++ b/iree/hal/cuda/graph_command_buffer.c
@@ -164,6 +164,18 @@ return iree_ok_status(); } +static void iree_hal_cuda_graph_command_buffer_begin_debug_group( + iree_hal_command_buffer_t* base_command_buffer, iree_string_view_t label, + iree_hal_label_color_t label_color, + const iree_hal_label_location_t* location) { + // TODO(benvanik): tracy event stack. +} + +static void iree_hal_cuda_graph_command_buffer_end_debug_group( + iree_hal_command_buffer_t* base_command_buffer) { + // TODO(benvanik): tracy event stack. +} + static iree_status_t iree_hal_cuda_graph_command_buffer_execution_barrier( iree_hal_command_buffer_t* base_command_buffer, iree_hal_execution_stage_t source_stage_mask, @@ -435,6 +447,9 @@ iree_hal_cuda_graph_command_buffer_allowed_categories, .begin = iree_hal_cuda_graph_command_buffer_begin, .end = iree_hal_cuda_graph_command_buffer_end, + .begin_debug_group = + iree_hal_cuda_graph_command_buffer_begin_debug_group, + .end_debug_group = iree_hal_cuda_graph_command_buffer_end_debug_group, .execution_barrier = iree_hal_cuda_graph_command_buffer_execution_barrier, .signal_event = iree_hal_cuda_graph_command_buffer_signal_event,
diff --git a/iree/hal/local/inline_command_buffer.c b/iree/hal/local/inline_command_buffer.c index 5cf10d1..3fd5692 100644 --- a/iree/hal/local/inline_command_buffer.c +++ b/iree/hal/local/inline_command_buffer.c
@@ -172,6 +172,22 @@ } //===----------------------------------------------------------------------===// +// iree_hal_inline_command_buffer_t debug utilities +//===----------------------------------------------------------------------===// + +static void iree_hal_inline_command_buffer_begin_debug_group( + iree_hal_command_buffer_t* base_command_buffer, iree_string_view_t label, + iree_hal_label_color_t label_color, + const iree_hal_label_location_t* location) { + // TODO(benvanik): tracy event stack. +} + +static void iree_hal_inline_command_buffer_end_debug_group( + iree_hal_command_buffer_t* base_command_buffer) { + // TODO(benvanik): tracy event stack. +} + +//===----------------------------------------------------------------------===// // iree_hal_command_buffer_execution_barrier //===----------------------------------------------------------------------===// @@ -484,6 +500,8 @@ .allowed_categories = iree_hal_inline_command_buffer_allowed_categories, .begin = iree_hal_inline_command_buffer_begin, .end = iree_hal_inline_command_buffer_end, + .begin_debug_group = iree_hal_inline_command_buffer_begin_debug_group, + .end_debug_group = iree_hal_inline_command_buffer_end_debug_group, .execution_barrier = iree_hal_inline_command_buffer_execution_barrier, .signal_event = iree_hal_inline_command_buffer_signal_event, .reset_event = iree_hal_inline_command_buffer_reset_event,
diff --git a/iree/hal/local/task_command_buffer.c b/iree/hal/local/task_command_buffer.c index b05c86c..2e78bd1 100644 --- a/iree/hal/local/task_command_buffer.c +++ b/iree/hal/local/task_command_buffer.c
@@ -379,6 +379,22 @@ } //===----------------------------------------------------------------------===// +// iree_hal_task_command_buffer_t debug utilities +//===----------------------------------------------------------------------===// + +static void iree_hal_task_command_buffer_begin_debug_group( + iree_hal_command_buffer_t* base_command_buffer, iree_string_view_t label, + iree_hal_label_color_t label_color, + const iree_hal_label_location_t* location) { + // TODO(benvanik): tracy event stack. +} + +static void iree_hal_task_command_buffer_end_debug_group( + iree_hal_command_buffer_t* base_command_buffer) { + // TODO(benvanik): tracy event stack. +} + +//===----------------------------------------------------------------------===// // iree_hal_command_buffer_execution_barrier //===----------------------------------------------------------------------===// @@ -896,6 +912,8 @@ .allowed_categories = iree_hal_task_command_buffer_allowed_categories, .begin = iree_hal_task_command_buffer_begin, .end = iree_hal_task_command_buffer_end, + .begin_debug_group = iree_hal_task_command_buffer_begin_debug_group, + .end_debug_group = iree_hal_task_command_buffer_end_debug_group, .execution_barrier = iree_hal_task_command_buffer_execution_barrier, .signal_event = iree_hal_task_command_buffer_signal_event, .reset_event = iree_hal_task_command_buffer_reset_event,
diff --git a/iree/hal/vulkan/direct_command_buffer.cc b/iree/hal/vulkan/direct_command_buffer.cc index 720ff14..96f8eed 100644 --- a/iree/hal/vulkan/direct_command_buffer.cc +++ b/iree/hal/vulkan/direct_command_buffer.cc
@@ -184,6 +184,13 @@ command_buffer->handle, &begin_info), "vkBeginCommandBuffer"); + IREE_VULKAN_TRACE_ZONE_BEGIN_EXTERNAL( + command_buffer->tracing_context, command_buffer->handle, + /*file_name=*/NULL, 0, + /*line=*/0, /*func_name=*/NULL, 0, + "iree_hal_vulkan_direct_command_buffer", + strlen("iree_hal_vulkan_direct_command_buffer")); + return iree_ok_status(); } @@ -192,6 +199,9 @@ iree_hal_vulkan_direct_command_buffer_t* command_buffer = iree_hal_vulkan_direct_command_buffer_cast(base_command_buffer); + IREE_VULKAN_TRACE_ZONE_END(command_buffer->tracing_context, + command_buffer->handle); + VK_RETURN_IF_ERROR( command_buffer->syms->vkEndCommandBuffer(command_buffer->handle), "vkEndCommandBuffer"); @@ -203,6 +213,49 @@ return iree_ok_status(); } +static void iree_hal_vulkan_direct_command_buffer_begin_debug_group( + iree_hal_command_buffer_t* base_command_buffer, iree_string_view_t label, + iree_hal_label_color_t label_color, + const iree_hal_label_location_t* location) { + iree_hal_vulkan_direct_command_buffer_t* command_buffer = + iree_hal_vulkan_direct_command_buffer_cast(base_command_buffer); + IREE_VULKAN_TRACE_ZONE_BEGIN_EXTERNAL( + command_buffer->tracing_context, command_buffer->handle, + location ? location->file.data : NULL, location ? location->file.size : 0, + location ? location->line : 0, /*func_name=*/NULL, 0, label.data, + label.size); + if (command_buffer->syms->vkCmdBeginDebugUtilsLabelEXT) { + char label_buffer[128]; + snprintf(label_buffer, sizeof(label_buffer), "%.*s", (int)label.size, + label.data); + VkDebugUtilsLabelEXT label_info = { + /*.sType=*/VK_STRUCTURE_TYPE_DEBUG_UTILS_LABEL_EXT, + /*.pNext=*/NULL, + /*.pLabelName=*/label_buffer, + /*.color=*/ + { + /*r=*/label_color.r / 255.0f, + /*g=*/label_color.g / 255.0f, + /*b=*/label_color.b / 255.0f, + /*a=*/label_color.a / 255.0f, + }, + }; + command_buffer->syms->vkCmdBeginDebugUtilsLabelEXT(command_buffer->handle, + &label_info); + } +} + +static void iree_hal_vulkan_direct_command_buffer_end_debug_group( + iree_hal_command_buffer_t* base_command_buffer) { + iree_hal_vulkan_direct_command_buffer_t* command_buffer = + iree_hal_vulkan_direct_command_buffer_cast(base_command_buffer); + if (command_buffer->syms->vkCmdEndDebugUtilsLabelEXT) { + command_buffer->syms->vkCmdEndDebugUtilsLabelEXT(command_buffer->handle); + } + IREE_VULKAN_TRACE_ZONE_END(command_buffer->tracing_context, + command_buffer->handle); +} + static VkPipelineStageFlags iree_hal_vulkan_convert_pipeline_stage_flags( iree_hal_execution_stage_t stage_mask) { VkPipelineStageFlags flags = 0; @@ -602,8 +655,8 @@ IREE_VULKAN_TRACE_ZONE_BEGIN_EXTERNAL( command_buffer->tracing_context, command_buffer->handle, source_location.file_name.data, source_location.file_name.size, - source_location.line, source_location.func_name.data, - source_location.func_name.size, NULL, 0); + source_location.line, /*func_name=*/NULL, 0, + source_location.func_name.data, source_location.func_name.size); }); // Get the compiled and linked pipeline for the specified entry point and @@ -638,8 +691,8 @@ IREE_VULKAN_TRACE_ZONE_BEGIN_EXTERNAL( command_buffer->tracing_context, command_buffer->handle, source_location.file_name.data, source_location.file_name.size, - source_location.line, source_location.func_name.data, - source_location.func_name.size, NULL, 0); + source_location.line, /*func_name=*/NULL, 0, + source_location.func_name.data, source_location.func_name.size); // Get the compiled and linked pipeline for the specified entry point and // bind it to the command buffer. @@ -671,6 +724,10 @@ iree_hal_vulkan_direct_command_buffer_allowed_categories, /*.begin=*/iree_hal_vulkan_direct_command_buffer_begin, /*.end=*/iree_hal_vulkan_direct_command_buffer_end, + /*.begin_debug_group=*/ + iree_hal_vulkan_direct_command_buffer_begin_debug_group, + /*.end_debug_group=*/ + iree_hal_vulkan_direct_command_buffer_end_debug_group, /*.execution_barrier=*/ iree_hal_vulkan_direct_command_buffer_execution_barrier, /*.signal_event=*/
diff --git a/iree/hal/vulkan/internal_vk_mem_alloc.cc b/iree/hal/vulkan/internal_vk_mem_alloc.cc index 9fa1b27..145afd9 100644 --- a/iree/hal/vulkan/internal_vk_mem_alloc.cc +++ b/iree/hal/vulkan/internal_vk_mem_alloc.cc
@@ -15,8 +15,9 @@ #define VMA_ASSERT IREE_DCHECK #define VMA_HEAVY_ASSERT IREE_DCHECK -// Use IREE_LOG for logging. -#ifndef NDEBUG +// NOTE: logging is disabled by default as unless you are debugging VMA itself +// the information is not useful and just slows things down. +#if 0 #define VMA_DEBUG_LOG(...) _IREE_LOG_INFO << __VA_ARGS__ #else #define VMA_DEBUG_LOG(...)
diff --git a/iree/modules/hal/exports.inl b/iree/modules/hal/exports.inl index f137bd7..8d8adc3 100644 --- a/iree/modules/hal/exports.inl +++ b/iree/modules/hal/exports.inl
@@ -41,12 +41,14 @@ EXPORT_FN("buffer_view.trace", iree_hal_module_buffer_view_trace, rCrD, v) EXPORT_FN("command_buffer.begin", iree_hal_module_command_buffer_begin, r, v) +EXPORT_FN("command_buffer.begin_debug_group", iree_hal_module_command_buffer_begin_debug_group, rr, v) EXPORT_FN("command_buffer.bind_descriptor_set", iree_hal_module_command_buffer_bind_descriptor_set, rrirCiD, v) EXPORT_FN("command_buffer.copy_buffer", iree_hal_module_command_buffer_copy_buffer, rririi, v) EXPORT_FN("command_buffer.create", iree_hal_module_command_buffer_create, rii, r) EXPORT_FN("command_buffer.dispatch", iree_hal_module_command_buffer_dispatch, rriiii, v) EXPORT_FN("command_buffer.dispatch.indirect", iree_hal_module_command_buffer_dispatch_indirect, rriri, v) EXPORT_FN("command_buffer.end", iree_hal_module_command_buffer_end, r, v) +EXPORT_FN("command_buffer.end_debug_group", iree_hal_module_command_buffer_end_debug_group, r, v) EXPORT_FN("command_buffer.execution_barrier", iree_hal_module_command_buffer_execution_barrier, riii, v) EXPORT_FN("command_buffer.fill_buffer", iree_hal_module_command_buffer_fill_buffer, rriii, v) EXPORT_FN("command_buffer.push_constants", iree_hal_module_command_buffer_push_constants, rriCiD, v)
diff --git a/iree/modules/hal/module.c b/iree/modules/hal/module.c index caf9a02..61ea37b 100644 --- a/iree/modules/hal/module.c +++ b/iree/modules/hal/module.c
@@ -556,6 +556,35 @@ return iree_hal_command_buffer_end(command_buffer); } +IREE_VM_ABI_EXPORT(iree_hal_module_command_buffer_begin_debug_group, // + iree_hal_module_state_t, // + rr, v) { + iree_hal_command_buffer_t* command_buffer = NULL; + IREE_RETURN_IF_ERROR( + iree_hal_command_buffer_check_deref(args->r0, &command_buffer)); + iree_vm_buffer_t* label = NULL; + IREE_RETURN_IF_ERROR(iree_vm_buffer_check_deref(args->r1, &label)); + iree_string_view_t label_str = iree_vm_buffer_as_string(label); + // TODO(benvanik): query from VM. + iree_hal_label_location_t location = { + .file = iree_string_view_empty(), + .line = 0, + }; + iree_hal_command_buffer_begin_debug_group( + command_buffer, label_str, iree_hal_label_color_unspecified(), &location); + return iree_ok_status(); +} + +IREE_VM_ABI_EXPORT(iree_hal_module_command_buffer_end_debug_group, // + iree_hal_module_state_t, // + r, v) { + iree_hal_command_buffer_t* command_buffer = NULL; + IREE_RETURN_IF_ERROR( + iree_hal_command_buffer_check_deref(args->r0, &command_buffer)); + iree_hal_command_buffer_end_debug_group(command_buffer); + return iree_ok_status(); +} + IREE_VM_ABI_EXPORT(iree_hal_module_command_buffer_execution_barrier, // iree_hal_module_state_t, // riii, v) {
diff --git a/iree/samples/static_library/CMakeLists.txt b/iree/samples/static_library/CMakeLists.txt new file mode 100644 index 0000000..a2f66e8 --- /dev/null +++ b/iree/samples/static_library/CMakeLists.txt
@@ -0,0 +1,80 @@ +# Copyright 2021 The IREE Authors +# +# Licensed under the Apache License v2.0 with LLVM Exceptions. +# See https://llvm.org/LICENSE.txt for license information. +# SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception + +if(NOT ${IREE_TARGET_BACKEND_DYLIB-LLVM-AOT} + OR NOT ${IREE_HAL_DRIVER_DYLIB} + OR NOT ${IREE_BUILD_COMPILER}) + return() +endif() + +# Set iree-translate binary +set(_TRANSLATE_TOOL_EXECUTABLE $<TARGET_FILE:iree_tools_iree-translate>) + +# Setup args for iree-transalte +set(_TRANSLATE_ARGS) +list(APPEND _TRANSLATE_ARGS "-iree-mlir-to-vm-bytecode-module") +list(APPEND _TRANSLATE_ARGS "-iree-hal-target-backends=dylib-llvm-aot") +list(APPEND _TRANSLATE_ARGS "-iree-llvm-static-library-output-path=simple_mul.o") +list(APPEND _TRANSLATE_ARGS "${CMAKE_CURRENT_SOURCE_DIR}/simple_mul.mlir") +list(APPEND _TRANSLATE_ARGS "-o") +list(APPEND _TRANSLATE_ARGS "simple_mul.vmfb") + +# Custom command for iree-translate to generate static library and bytecode +add_custom_command( + OUTPUT + ${CMAKE_CURRENT_BINARY_DIR}/simple_mul.h + ${CMAKE_CURRENT_BINARY_DIR}/simple_mul.o + ${CMAKE_CURRENT_BINARY_DIR}/simple_mul.vmfb + COMMAND ${_TRANSLATE_TOOL_EXECUTABLE} ${_TRANSLATE_ARGS} + DEPENDS ${_TRANSLATE_TOOL_EXECUTABLE} "simple_mul.mlir" +) + +# Tell cmake about the simple_mul library so it will link it +add_library(simple_mul + STATIC + ${CMAKE_CURRENT_BINARY_DIR}/simple_mul.o) + +SET_TARGET_PROPERTIES( + simple_mul + PROPERTIES + LINKER_LANGUAGE C) + +# Note: If you're cross compiling the simple_mul for a different backend, you'll +# need to run iree-translate manually to produce the desired '.vmfb' and '.h/.o' +# files. Substitute the 'simple_mul' dependency in iree_cc_binary() below with +# your own static library and the `simple_mul.vmfb` in the iree_c_embed_data() +# rule. You can use paths to files, i.e. 'path/to/generated/output.vmfb' in +# place of CMake targets. + +# Generate the embed data with the bytecode module +iree_c_embed_data( + NAME + simple_mul_c + IDENTIFIER + iree_samples_static_library_simple_mul + GENERATED_SRCS + simple_mul.vmfb + C_FILE_OUTPUT + simple_mul_c.c + H_FILE_OUTPUT + simple_mul_c.h + FLATTEN + PUBLIC +) + +iree_cc_binary( +NAME + static_library_demo +SRCS + "static_library_demo.c" +DEPS + ::simple_mul_c + iree::runtime + iree::hal::local::loaders::static_library_loader + iree::hal::local::task_driver + iree::task::api + simple_mul +)
diff --git a/iree/samples/static_library/README.md b/iree/samples/static_library/README.md new file mode 100644 index 0000000..b989f0a --- /dev/null +++ b/iree/samples/static_library/README.md
@@ -0,0 +1,79 @@ +# IREE "Static Library" sample + +This sample shows how to: +1. Produce a static library and bytecode module with IREE's compiler +2. Compile the static library into a program using the `static_library_loader` +3. Run the demo with the module using functions exported by the static library + +The model compiled into the static library exports a single function +`simple_mul` that returns the multiplication of two tensors: + +```mlir +func @simple_mul(%arg0: tensor<4xf32>, %arg1: tensor<4xf32>) -> tensor<4xf32> +{ + %0 = "mhlo.multiply"(%arg0, %arg1) {name = "mul.1"} : (tensor<4xf32>, + tensor<4xf32>) -> tensor<4xf32> + return %0 : tensor<4xf32> +} +``` + +## Background + +IREE's `static_library_loader` allows applications to inject a set of static +libraries that can be resolved at runtime by name. This can be particularly +useful on "bare metal" or embedded systems running IREE that lack operating +systems or the ability to load shared libraries in binaries. + +When static library output is enabled, `iree-translate` produces a separate +static library to compile into the target program. At runtime bytecode module +instructs the VM which static libraries to load exported functions from the +model. + +## Instructions +_Note: run the following commands from IREE's github repo root._ + +1. Configure CMake for building the static library then demo. You'll need to set +the flags building samples, the compiler, and the `dylib-llvm-aot` +driver/backend. See +[here](https://google.github.io/iree/building-from-source/getting-started/) +for general instructions on building using CMake): + + ```shell + cmake -B ../iree-build/ + -DIREE_BUILD_SAMPLES=ON \ + -DIREE_TARGET_BACKENDS_TO_BUILD=DYLIB-LLVM-AOT \ + -DIREE_HAL_DRIVERS_TO_BUILD=DYLIB \ + -DIREE_BUILD_COMPILER=ON \ + -DCMAKE_BUILD_TYPE=RelWithDebInfo . + ``` + +2. Build the `static_library_demo` CMake target to create the static demo. This +target has several dependencies that will translate `simple_mul.mlir` into a +static library (`simple_mul.h` & `simple_mul.c`) as well as a bytecode module +(`simple_mul.vmfb`) which are finally built into the demo binary: + + ```shell + cmake --build ../iree-build/ --target iree_samples_static_library_demo + ``` + +3. Run the sample binary: + + ```shell + ../iree-build/iree/samples/static_library/static_library_demo + + # Output: static_library_run passed + ``` + +### Changing compilation options + +The steps above build both the compiler for the host (machine doing the +compiling) and the demo for the target using same options as the host machine. +If you wish to target a different deployment other than the host, you'll need to +compile the library and demo with different options. + +For example, see +[documentation](https://google.github.io/iree/building-from-source/android/) +on cross compiling on Android. + +Note: separating the target from the host will require modifying dependencies in +the demos `CMakeLists.txt`. See included comments for more info.
diff --git a/iree/samples/static_library/simple_mul.mlir b/iree/samples/static_library/simple_mul.mlir new file mode 100644 index 0000000..2b480bd --- /dev/null +++ b/iree/samples/static_library/simple_mul.mlir
@@ -0,0 +1,5 @@ +func @simple_mul(%arg0: tensor<4xf32>, %arg1: tensor<4xf32>) -> tensor<4xf32> +{ + %0 = "std.mulf"(%arg0, %arg1) {name = "mul.1"} : (tensor<4xf32>, tensor<4xf32>) -> tensor<4xf32> + return %0 : tensor<4xf32> +}
diff --git a/iree/samples/static_library/static_library_demo.c b/iree/samples/static_library/static_library_demo.c new file mode 100644 index 0000000..22e8127 --- /dev/null +++ b/iree/samples/static_library/static_library_demo.c
@@ -0,0 +1,222 @@ +// Copyright 2021 The IREE Authors +// +// Licensed under the Apache License v2.0 with LLVM Exceptions. +// See https://llvm.org/LICENSE.txt for license information. +// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception + +// A example of static library loading in IREE. See the README.md for more info. +// Note: this demo requires artifacts from iree-translate before it will run. + +#include <stdio.h> + +#include "iree/hal/local/loaders/static_library_loader.h" +#include "iree/hal/local/task_device.h" +#include "iree/modules/hal/module.h" +#include "iree/runtime/api.h" +#include "iree/samples/static_library/simple_mul_c.h" +#include "iree/task/api.h" +#include "iree/vm/bytecode_module.h" + +// Compiled static library module here to avoid IO: +#include "iree/samples/static_library/simple_mul.h" + +// A function to create the HAL device from the different backend targets. +// The HAL device is returned based on the implementation, and it must be +// released by the caller. +iree_status_t create_device_with_static_loader(iree_hal_device_t** device) { + iree_status_t status = iree_ok_status(); + // Set paramters for the device created in the next step. + iree_hal_task_device_params_t params; + iree_hal_task_device_params_initialize(¶ms); + + // Load the statically embedded library + const iree_hal_executable_library_header_t** static_library = + simple_mul_dispatch_0_library_query( + IREE_HAL_EXECUTABLE_LIBRARY_LATEST_VERSION, /*reserved=*/NULL); + const iree_hal_executable_library_header_t** libraries[1] = {static_library}; + + iree_hal_executable_loader_t* library_loader = NULL; + if (iree_status_is_ok(status)) { + status = iree_hal_static_library_loader_create( + IREE_ARRAYSIZE(libraries), libraries, + iree_hal_executable_import_provider_null(), iree_allocator_system(), + &library_loader); + } + + iree_task_executor_t* executor = NULL; + if (iree_status_is_ok(status)) { + iree_task_executor_create_from_flags(iree_allocator_system(), &executor); + } + + // Create the device and release the executor and loader afterwards. + if (iree_status_is_ok(status)) { + iree_hal_task_device_create(iree_make_cstring_view("dylib"), ¶ms, + executor, 1, &library_loader, + iree_allocator_system(), device); + } + iree_task_executor_release(executor); + iree_hal_executable_loader_release(library_loader); + + return status; +} + +iree_status_t Run() { + iree_status_t status = iree_ok_status(); + + // Instance configuration (this should be shared across sessions). + iree_runtime_instance_options_t instance_options; + iree_runtime_instance_options_initialize(IREE_API_VERSION_LATEST, + &instance_options); + iree_runtime_instance_options_use_all_available_drivers(&instance_options); + iree_runtime_instance_t* instance = NULL; + + if (iree_status_is_ok(status)) { + status = iree_runtime_instance_create(&instance_options, + iree_allocator_system(), &instance); + } + + // Create dylib device with static loader. + iree_hal_device_t* device = NULL; + if (iree_status_is_ok(status)) { + status = create_device_with_static_loader(&device); + } + iree_vm_module_t* hal_module = NULL; + if (iree_status_is_ok(status)) { + status = + iree_hal_module_create(device, iree_allocator_system(), &hal_module); + } + + // Session configuration (one per loaded module to hold module state). + iree_runtime_session_options_t session_options; + iree_runtime_session_options_initialize(&session_options); + iree_runtime_session_t* session = NULL; + if (iree_status_is_ok(status)) { + status = iree_runtime_session_create_with_device( + instance, &session_options, device, + iree_runtime_instance_host_allocator(instance), &session); + } + + // Load bytecode module from the embedded data. Append to the session. + const struct iree_file_toc_t* module_file_toc = + iree_samples_static_library_simple_mul_create(); + iree_const_byte_span_t module_data = + iree_make_const_byte_span(module_file_toc->data, module_file_toc->size); + iree_vm_module_t* bytecode_module = NULL; + if (iree_status_is_ok(status)) { + status = iree_vm_bytecode_module_create(module_data, iree_allocator_null(), + iree_allocator_system(), + &bytecode_module); + } + if (iree_status_is_ok(status)) { + status = iree_runtime_session_append_module(session, bytecode_module); + } + + // Lookup the entry point function call. + const char kMainFunctionName[] = "module.simple_mul"; + iree_runtime_call_t call; + memset(&call, 0, sizeof(call)); + if (iree_status_is_ok(status)) { + status = iree_runtime_call_initialize_by_name( + session, iree_make_cstring_view(kMainFunctionName), &call); + } + + // Populate initial values for 4 * 2 = 8. + const int kElementCount = 4; + iree_hal_dim_t shape[1] = {kElementCount}; + iree_hal_buffer_view_t* arg0_buffer_view = NULL; + iree_hal_buffer_view_t* arg1_buffer_view = NULL; + float kFloat4[] = {4.0f, 4.0f, 4.0f, 4.0f}; + float kFloat2[] = {2.0f, 2.0f, 2.0f, 2.0f}; + + iree_hal_memory_type_t input_memory_type = + IREE_HAL_MEMORY_TYPE_HOST_LOCAL | IREE_HAL_MEMORY_TYPE_DEVICE_VISIBLE; + if (iree_status_is_ok(status)) { + iree_hal_buffer_view_clone_heap_buffer( + iree_hal_device_allocator(device), shape, IREE_ARRAYSIZE(shape), + IREE_HAL_ELEMENT_TYPE_FLOAT_32, input_memory_type, + IREE_HAL_BUFFER_USAGE_ALL, + iree_make_const_byte_span((void*)kFloat4, + sizeof(float) * kElementCount), + &arg0_buffer_view); + } + if (iree_status_is_ok(status)) { + iree_hal_buffer_view_clone_heap_buffer( + iree_hal_device_allocator(device), shape, IREE_ARRAYSIZE(shape), + IREE_HAL_ELEMENT_TYPE_FLOAT_32, input_memory_type, + IREE_HAL_BUFFER_USAGE_ALL, + iree_make_const_byte_span((void*)kFloat2, + sizeof(float) * kElementCount), + &arg1_buffer_view); + } + + // Queue buffer views for input. + if (iree_status_is_ok(status)) { + status = + iree_runtime_call_inputs_push_back_buffer_view(&call, arg0_buffer_view); + } + iree_hal_buffer_view_release(arg0_buffer_view); + + if (iree_status_is_ok(status)) { + status = + iree_runtime_call_inputs_push_back_buffer_view(&call, arg1_buffer_view); + } + iree_hal_buffer_view_release(arg1_buffer_view); + + // Invoke call. + if (iree_status_is_ok(status)) { + status = iree_runtime_call_invoke(&call, /*flags=*/0); + } + + // Retreive output buffer view with results from the invocation. + iree_hal_buffer_view_t* ret_buffer_view = NULL; + if (iree_status_is_ok(status)) { + status = iree_runtime_call_outputs_pop_front_buffer_view(&call, + &ret_buffer_view); + } + + // Read back the results and ensure we got the right values. + iree_hal_buffer_mapping_t mapped_memory = {0}; + if (iree_status_is_ok(status)) { + status = iree_hal_buffer_map_range( + iree_hal_buffer_view_buffer(ret_buffer_view), + IREE_HAL_MEMORY_ACCESS_READ, 0, IREE_WHOLE_BUFFER, &mapped_memory); + } + if (iree_status_is_ok(status)) { + if (mapped_memory.contents.data_length / sizeof(float) != kElementCount) { + status = iree_make_status(IREE_STATUS_UNKNOWN, + "result does not match element count "); + } + } + if (iree_status_is_ok(status)) { + const float* data = (const float*)mapped_memory.contents.data; + for (iree_host_size_t i = 0; + i < mapped_memory.contents.data_length / sizeof(float); ++i) { + if (data[i] != 8.0f) { + status = iree_make_status(IREE_STATUS_UNKNOWN, "result mismatches"); + } + } + } + + // Cleanup call and buffers. + iree_hal_buffer_unmap_range(&mapped_memory); + iree_hal_buffer_view_release(ret_buffer_view); + iree_runtime_call_deinitialize(&call); + + // Cleanup session and instance. + iree_hal_device_release(device); + iree_runtime_session_release(session); + iree_runtime_instance_release(instance); + + return status; +} + +int main() { + const iree_status_t result = Run(); + if (!iree_status_is_ok(result)) { + iree_status_fprint(stderr, result); + iree_status_free(result); + return -1; + } + printf("static_library_run passed\n"); + return 0; +}
diff --git a/iree/task/task.c b/iree/task/task.c index 160b6c1..51fc703 100644 --- a/iree/task/task.c +++ b/iree/task/task.c
@@ -824,10 +824,6 @@ tiles_per_reservation, iree_memory_order_relaxed); while (tile_base < tile_count) { - const uint32_t next_tile_base = iree_atomic_fetch_add_int32( - &shared_state->tile_index, tiles_per_reservation, - iree_memory_order_relaxed); - const uint32_t tile_range = iree_min(tile_base + tiles_per_reservation, tile_count); for (uint32_t tile_index = tile_base; tile_index < tile_range; @@ -865,7 +861,9 @@ } } - tile_base = next_tile_base; + tile_base = iree_atomic_fetch_add_int32(&shared_state->tile_index, + tiles_per_reservation, + iree_memory_order_relaxed); } // Push aggregate statistics up to the dispatch.
diff --git a/iree/vm/stack.c b/iree/vm/stack.c index e7005fb..eadee60 100644 --- a/iree/vm/stack.c +++ b/iree/vm/stack.c
@@ -413,14 +413,18 @@ stack->frame_storage_size = new_top; stack->top = frame_header; -#if IREE_TRACING_FEATURES & IREE_TRACING_FEATURE_INSTRUMENTATION - // TODO(benvanik): cache source location and query from module. - iree_string_view_t function_name = iree_vm_function_name(function); - IREE_TRACE_ZONE_BEGIN_NAMED_DYNAMIC(z0, function_name.data, - function_name.size); - callee_frame->trace_zone = z0; - IREE_TRACE_ZONE_APPEND_VALUE(z0, frame_size); -#endif // IREE_TRACING_FEATURES & IREE_TRACING_FEATURE_INSTRUMENTATION + IREE_TRACE({ + if (frame_type != IREE_VM_STACK_FRAME_NATIVE) { + // TODO(benvanik): cache source location and query from module. + iree_string_view_t function_name = iree_vm_function_name(function); + IREE_TRACE_ZONE_BEGIN_NAMED_DYNAMIC(z0, function_name.data, + function_name.size); + callee_frame->trace_zone = z0; + if (frame_size) { + IREE_TRACE_ZONE_APPEND_VALUE(z0, frame_size); + } + } + }); if (out_callee_frame) *out_callee_frame = callee_frame; return iree_ok_status(); @@ -438,7 +442,11 @@ stack->top->frame_cleanup_fn(&stack->top->frame); } - IREE_TRACE_ZONE_END(stack->top->frame.trace_zone); + IREE_TRACE({ + if (stack->top->frame.trace_zone) { + IREE_TRACE_ZONE_END(stack->top->frame.trace_zone); + } + }); // Restore the frame pointer to the caller. stack->frame_storage_size -= stack->top->frame_size;
diff --git a/scripts/ir_to_markdown.py b/scripts/ir_to_markdown.py new file mode 100644 index 0000000..952797c --- /dev/null +++ b/scripts/ir_to_markdown.py
@@ -0,0 +1,103 @@ +#!/usr/bin/env python3 + +# Copyright 2021 The IREE Authors +# +# Licensed under the Apache License v2.0 with LLVM Exceptions. +# See https://llvm.org/LICENSE.txt for license information. +# SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception +"""Converts a dump of IR text to a markdown file with formatting. + +Example usage: + # Get a dump of IR from a compiler tool: + $ iree-opt \ + -iree-transformation-pipeline \ + -iree-hal-target-backends=vmvx \ + -mlir-disable-threading \ + -print-ir-after-all \ + -print-ir-after-change \ + -mlir-elide-elementsattrs-if-larger=8 \ + $PWD/iree/samples/models/simple_abs.mlir \ + 2> /tmp/simple_abs_vmvx_pipeline.mlir + > /dev/null + + # Convert the IR dump to markdown: + $ python3 ir_to_markdown.py \ + /tmp/simple_abs_vmvx_pipeline.mlir \ + -o /tmp/simple_abs_vmvx_pipeline.md +""" + +import argparse +import re + +MLIR_START_SEQUENCE = "// -----//" +MLIR_END_SEQUENCE = "//----- //" + + +def parse_arguments(): + """Parses command line arguments.""" + + parser = argparse.ArgumentParser() + parser.add_argument('input_file_path', + type=str, + nargs='?', + metavar="<input_file_path>", + help='Input IR dump (.mlir from -print-ir-after-all)') + parser.add_argument('-o,', + '--output', + type=str, + required=True, + metavar="<output>", + help='Output file path (e.g. translation_ir.md)') + # TODO(scotttodd): flags for original IR path and compilation command line + # .md could then show original IR + flags -> output + # TODO(scotttodd): flag for markdown flavor (mkdocs, github, etc.) + # TODO(scotttodd): flag for diff view (correlate IR before and IR after)? + + return parser.parse_args() + + +def main(args): + input_file_path = args.input_file_path + output_file_path = args.output + print("Converting input file '%s'" % (input_file_path)) + print(" into output file '%s'" % (output_file_path)) + + with open(input_file_path, "r") as input_file: + with open(output_file_path, "w") as output_file: + + # Iterate line by line through the input file, collecting text into + # blocks and writing them into the output file with markdown formatting + # as we go. + # + # Note: we could parse through and find/replace within the file using + # regex (or sed), but iterating this way is easier to understand and + # uses a predictable amount of memory. + + current_block_lines = [] + dump_after_regex = re.compile(MLIR_START_SEQUENCE + "\s(.*)\s" + + MLIR_END_SEQUENCE) + + def finish_block(): + nonlocal current_block_lines + if len(current_block_lines) != 0: + current_block_lines.append("```\n\n") + output_file.writelines(current_block_lines) + current_block_lines = [] + + for input_line in input_file: + if input_line == "\n": + continue + + if input_line.startswith(MLIR_START_SEQUENCE): + finish_block() + header_text = dump_after_regex.match(input_line).group(1) + current_block_lines.append("### " + header_text + "\n\n") + current_block_lines.append("```mlir\n") + else: + current_block_lines.append(input_line) + + finish_block() + + +if __name__ == '__main__': + main(parse_arguments())
diff --git a/third_party/tracy b/third_party/tracy index f03b677..50f7deb 160000 --- a/third_party/tracy +++ b/third_party/tracy
@@ -1 +1 @@ -Subproject commit f03b677ffa0fd96fcf859c32e79b740fac7dd59e +Subproject commit 50f7deb1a389bd3785c12fbe0be74128343f11f7