From 5bd63d62324a289cdd40ec00c8ce6200cb27116a Mon Sep 17 00:00:00 2001 From: "Yukihiro \"Matz\" Matsumoto" Date: Wed, 26 Jun 2024 08:24:48 +0900 Subject: [PATCH] array.c: replace sort! method implementation - use heap sort (O(1)) instead of merge sort (O(n)) for better space complexity. - method implemented in C for better performance As a result, simple sorting now consumes far less memory and is faster. Since it's implemented in C, fiber context switching is not allowed from comparison, but we consider the risk is minimal (no one switches context in the comparison, right?) --- mrblib/array.rb | 72 ------------------------------------------------- src/array.c | 71 ++++++++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 71 insertions(+), 72 deletions(-) diff --git a/mrblib/array.rb b/mrblib/array.rb index 76c137421..5ce61928f 100644 --- a/mrblib/array.rb +++ b/mrblib/array.rb @@ -198,78 +198,6 @@ class Array ret end - ## - # call-seq: - # array.sort! -> self - # array.sort! {|a, b| ... } -> self - # - # Sort all elements and replace +self+ with these - # elements. - def sort!(&block) - stack = [ [ 0, self.size - 1 ] ] - until stack.empty? - left, mid, right = stack.pop - if right == nil - right = mid - # sort self[left..right] - if left < right - if left + 1 == right - lval = self[left] - rval = self[right] - cmp = if block then block.call(lval,rval) else lval <=> rval end - if cmp.nil? - raise ArgumentError, "comparison of #{lval.inspect} and #{rval.inspect} failed" - end - if cmp > 0 - self[left] = rval - self[right] = lval - end - else - mid = ((left + right + 1) / 2).floor - stack.push [ left, mid, right ] - stack.push [ mid, right ] - stack.push [ left, (mid - 1) ] if left < mid - 1 - end - end - else - lary = self[left, mid - left] - lsize = lary.size - - # The entity sharing between lary and self may cause a large memory - # copy operation in the merge loop below. This harmless operation - # cancels the sharing and provides a huge performance gain. - lary[0] = lary[0] - - # merge - lidx = 0 - ridx = mid - (left..right).each { |i| - if lidx >= lsize - break - elsif ridx > right - self[i, lsize - lidx] = lary[lidx, lsize - lidx] - break - else - lval = lary[lidx] - rval = self[ridx] - cmp = if block then block.call(lval,rval) else lval <=> rval end - if cmp.nil? - raise ArgumentError, "comparison of #{lval.inspect} and #{rval.inspect} failed" - end - if cmp <= 0 - self[i] = lval - lidx += 1 - else - self[i] = rval - ridx += 1 - end - end - } - end - end - self - end - ## # call-seq: # array.sort -> new_array diff --git a/src/array.c b/src/array.c index 74d4f4e2e..eb3474842 100644 --- a/src/array.c +++ b/src/array.c @@ -1409,6 +1409,76 @@ mrb_ary_delete(mrb_state *mrb, mrb_value self) return ret; } +static mrb_bool +sort_cmp(mrb_state *mrb, mrb_value *p, mrb_int a, mrb_int b, mrb_value blk) +{ + mrb_int cmp; + + if (mrb_nil_p(blk)) { + cmp = mrb_cmp(mrb, p[a], p[b]); + } + else { + mrb_value c = mrb_funcall_id(mrb, blk, MRB_SYM(call), 2, p[a], p[b]); + if (mrb_nil_p(c) || !mrb_fixnum_p(c)) { + mrb_raisef(mrb, E_ARGUMENT_ERROR, "comparison of %!v and %!v failed", p[a], p[b]); + } + cmp = mrb_fixnum(c); + } + return cmp > 0; +} + +static void +heapify(mrb_state *mrb, mrb_value *a, mrb_int index, mrb_int size, mrb_value blk) +{ + mrb_int max = index; + mrb_int left_index = 2 * index + 1; + mrb_int right_index = left_index + 1; + if (left_index < size && sort_cmp(mrb, a, left_index, max, blk)) { + max = left_index; + } + if (right_index < size && sort_cmp(mrb, a, right_index, max, blk)) { + max = right_index; + } + if (max != index) { + mrb_value tmp = a[max]; + a[max] = a[index]; + a[index] = tmp; + heapify(mrb, a, max, size, blk); + } +} + +/* + * call-seq: + * array.sort! -> self + * array.sort! {|a, b| ... } -> self + * + * Sort all elements and replace +self+ with these + * elements. + */ +static mrb_value +mrb_ary_sort_bang(mrb_state *mrb, mrb_value ary) +{ + mrb_value blk; + + mrb_int n = RARRAY_LEN(ary); + if (n < 2) return ary; + + ary_modify(mrb, mrb_ary_ptr(ary)); + mrb_get_args(mrb, "&", &blk); + + mrb_value *a = RARRAY_PTR(ary); + for (mrb_int i = n / 2 - 1; i > -1; i--) { + heapify(mrb, a, i, n, blk); + } + for (mrb_int i = n - 1; i > 0; i--) { + mrb_value tmp = a[0]; + a[0] = a[i]; + a[i] = tmp; + heapify(mrb, a, 0, i, blk); + } + return ary; +} + void mrb_init_array(mrb_state *mrb) { @@ -1446,6 +1516,7 @@ mrb_init_array(mrb_state *mrb) mrb_define_method_id(mrb, a, MRB_SYM(unshift), mrb_ary_unshift_m, MRB_ARGS_ANY()); /* 15.2.12.5.30 */ mrb_define_method_id(mrb, a, MRB_SYM(to_s), mrb_ary_to_s, MRB_ARGS_NONE()); mrb_define_method_id(mrb, a, MRB_SYM(inspect), mrb_ary_to_s, MRB_ARGS_NONE()); + mrb_define_method_id(mrb, a, MRB_SYM_B(sort), mrb_ary_sort_bang, MRB_ARGS_NONE()); mrb_define_method_id(mrb, a, MRB_SYM(__ary_eq), mrb_ary_eq, MRB_ARGS_REQ(1)); mrb_define_method_id(mrb, a, MRB_SYM(__ary_cmp), mrb_ary_cmp, MRB_ARGS_REQ(1));