二つのソート済み配列の中央値を二分探索で求める方法

問題の理解

二つのソート済み配列が与えられた場合、全体の中央値を効率的に見つける必要があります。単純な方法では両方の配列をマージしてから中央値を計算できますが、これではO(m+n)の時間計算量が必要です。より効率的な解法として、二分探索を用いることでO(log(min(m,n)))の時間計算量で解くことができます。

アルゴリズムの考え方

二つの配列から、左半分の要素数をk=(m+n+1)/2にします。ここで、最初の配列からi個の要素を選び、2番目の配列からk-i個の要素を選びます。選んだ要素が左半分を形成するためには、以下の条件を満たす必要があります:

  • 最初の配列のi番目の要素が、2番目の配列の(k-i-1)番目の要素以上であること
  • 2番目の配列の(k-i)番目の要素が、最初の配列の(i-1)番目の要素以下であること

この条件を満たすiを見つけるために二分探索を使用します。

実装例

Pythonでの実装

class MedianFinder:
    def find_median(self, arr1, arr2):
        len1, len2 = len(arr1), len(arr2)
        
        # 最初の配列が短い方になるようにする
        if len1 > len2:
            return self.find_median(arr2, arr1)
            
        total = len1 + len2
        half = (total + 1) // 2
        
        low, high = 0, len1
        
        while low < high:
            partition1 = (low + high) // 2
            partition2 = half - partition1
            
            # 左側の最大値と右側の最小値を比較
            left_max1 = float('-inf') if partition1 == 0 else arr1[partition1 - 1]
            left_max2 = float('-inf') if partition2 == 0 else arr2[partition2 - 1]
            right_min1 = float('inf') if partition1 == len1 else arr1[partition1]
            right_min2 = float('inf') if partition2 == len2 else arr2[partition2]
            
            if left_max1 <= right_min2 and left_max2 <= right_min1:
                # 正しい分割が見つかった
                if total % 2 == 1:
                    return max(left_max1, left_max2)
                else:
                    return (max(left_max1, left_max2) + min(right_min1, right_min2)) / 2.0
            elif left_max1 > right_min2:
                # 分割点を左に移動
                high = partition1 - 1
            else:
                # 分割点を右に移動
                low = partition1 + 1
                
        # ここに到達することはないはず
        return 0.0

C++での実装

class MedianCalculator {
public:
    double findMedianSortedArrays(vector<int>& vec1, vector<int>& vec2) {
        int size1 = vec1.size();
        int size2 = vec2.size();
        
        // 最初の配列が短い方になるようにする
        if (size1 > size2) {
            return findMedianSortedArrays(vec2, vec1);
        }
        
        int total = size1 + size2;
        int half = (total + 1) / 2;
        
        int left = 0;
        int right = size1;
        
        while (left <= right) {
            int partition1 = (left + right) / 2;
            int partition2 = half - partition1;
            
            int max_left1 = (partition1 == 0) ? INT_MIN : vec1[partition1 - 1];
            int max_left2 = (partition2 == 0) ? INT_MIN : vec2[partition2 - 1];
            int min_right1 = (partition1 == size1) ? INT_MAX : vec1[partition1];
            int min_right2 = (partition2 == size2) ? INT_MAX : vec2[partition2];
            
            if (max_left1 <= min_right2 && max_left2 <= min_right1) {
                if (total % 2 == 1) {
                    return max(max_left1, max_left2);
                } else {
                    return (max(max_left1, max_left2) + min(min_right1, min_right2)) / 2.0;
                }
            } else if (max_left1 > min_right2) {
                right = partition1 - 1;
            } else {
                left = partition1 + 1;
            }
        }
        
        return 0.0;
    }
};

Javaでの実装

public class MedianSolution {
    public double findMedian(int[] array1, int[] array2) {
        int len1 = array1.length;
        int len2 = array2.length;
        
        // 最初の配列が短い方になるようにする
        if (len1 > len2) {
            return findMedian(array2, array1);
        }
        
        int totalLength = len1 + len2;
        int halfSize = (totalLength + 1) / 2;
        
        int low = 0;
        int high = len1;
        
        while (low <= high) {
            int partition1 = (low + high) / 2;
            int partition2 = halfSize - partition1;
            
            int leftMax1 = (partition1 == 0) ? Integer.MIN_VALUE : array1[partition1 - 1];
            int leftMax2 = (partition2 == 0) ? Integer.MIN_VALUE : array2[partition2 - 1];
            int rightMin1 = (partition1 == len1) ? Integer.MAX_VALUE : array1[partition1];
            int rightMin2 = (partition2 == len2) ? Integer.MAX_VALUE : array2[partition2];
            
            if (leftMax1 <= rightMin2 && leftMax2 <= rightMin1) {
                if (totalLength % 2 == 1) {
                    return Math.max(leftMax1, leftMax2);
                } else {
                    return (Math.max(leftMax1, leftMax2) + Math.min(rightMin1, rightMin2)) / 2.0;
                }
            } else if (leftMax1 > rightMin2) {
                high = partition1 - 1;
            } else {
                low = partition1 + 1;
            }
        }
        
        return 0.0;
    }
}

アルゴリズムの複雑さ

このアルゴリズムの時間計算量はO(log(min(m,n)))です。これは、二分探索を短い配列の長さに対して行うためです。空間計算量はO(1)で、追加の空間をほとんど使用しないためです。

例えば、配列A = [-1, 1, 3, 5, 7, 9]と配列B = [2, 4, 6, 8, 10, 12, 14, 16]が与えられた場合、中央値は6.5となります。

タグ: アルゴリズム 二分探索 中央値 ソート済み配列

7月31日 01:06 投稿