<?xml version="1.0" encoding="utf-8"?> 
<feed xmlns="http://www.w3.org/2005/Atom" xml:lang="en-us">
    <generator uri="https://gohugo.io/" version="0.141.0">Hugo</generator><title type="html"><![CDATA[Binary-Search on Blog]]></title>
    
    
    
            <link href="https://blog.scientific-python.org/tags/binary-search/" rel="alternate" type="text/html" title="html" />
            <link href="https://blog.scientific-python.org/tags/binary-search/atom.xml" rel="self" type="application/atom" title="atom" />
    <updated>2026-09-30T00:53:53+00:00</updated>
    
    
    
    
        <id>https://blog.scientific-python.org/tags/binary-search/</id>
    
        
        <entry>
            <title type="html"><![CDATA[Making np.searchsorted up to 25× Faster in NumPy 2.5]]></title>
            <link href="https://blog.scientific-python.org/numpy/searchsorted/?utm_source=atom_feed" rel="alternate" type="text/html" />
            
                <link href="https://blog.scientific-python.org/numpy/fellowship-program-2025-retrospective/?utm_source=atom_feed" rel="related" type="text/html" title="A Year of Typing: My NumPy Fellowship Retrospective" />
                <link href="https://blog.scientific-python.org/numpy/fellowship-program-2025/?utm_source=atom_feed" rel="related" type="text/html" title="NumPy&#39;s Second Developer in Residence: Joren Hammudoglu" />
                <link href="https://blog.scientific-python.org/numpy/numpy2/?utm_source=atom_feed" rel="related" type="text/html" title="NumPy 2.0: an evolutionary milestone" />
                <link href="https://blog.scientific-python.org/numpy/numpy-rng/?utm_source=atom_feed" rel="related" type="text/html" title="Best Practices for Using NumPy&#39;s Random Number Generators" />
                <link href="https://blog.scientific-python.org/numpy/fellowship-program/?utm_source=atom_feed" rel="related" type="text/html" title="NumPy&#39;s first Developer in Residence: Sayed Adel" />
            
                <id>https://blog.scientific-python.org/numpy/searchsorted/</id>
            
            
            <published>2026-09-29T00:00:00+00:00</published>
            <updated>2026-09-29T00:00:00+00:00</updated>
            
            
            <content type="html"><![CDATA[<blockquote>How vectorizing independent binary searches and reducing per-query state can make np.searchsorted up to 25× faster.</blockquote><p><code>np.searchsorted</code> is NumPy&rsquo;s implementation of the binary search algorithm. This is one of the fundamental search algorithms and is used in the Python scientific ecosystem for functionality such as histogram computation and interval lookups. Any optimization benefits libraries such as SciPy and scikit-learn, as well as the broader Python scientific ecosystem.</p>
<p>Several case studies, such as <a href="https://curiouscoding.nl/posts/binsearch/">Binary search variants and the effects of batching</a> and <a href="https://en.algorithmica.org/hpc/data-structures/binary-search/">Algorithmica&rsquo;s Binary Search case study</a> explore techniques such as branch elimination, batching, and cache-friendly data layouts to binary search performance. In this post, we explore how those ideas can be expressed using NumPy&rsquo;s vectorized primitives.</p>
<p>We will derive a vectorized formulation that outperforms NumPy 2.4&rsquo;s searchsorted implementation, and then port the resulting algorithm back into NumPy. The change was included in NumPy 2.5, achieving up to a 25× speedup in our benchmarks.</p>
<p><code>searchsorted</code> is also part of the <a href="https://data-apis.org/array-api/latest/API_specification/generated/array_api.searchsorted.html#array_api.searchsorted">Python Array API Standard</a>. This allows us to compare how different array libraries implement the same operation and exploit parallelism.</p>
<h2 id="problem-definition">Problem definition<a class="headerlink" href="#problem-definition" title="Link to this heading">#</a></h2>
<p>Given a static sorted array and a sequence of query keys, find the insertion position of each key in the array.</p>
<p>A classic pure-Python implementation runs one binary search per key:</p>


<div class="highlight">
  <pre class="chroma"><code><span class="line"><span class="cl"><span class="k">def</span> <span class="nf">searchsorted_py</span><span class="p">(</span><span class="n">a</span><span class="p">,</span> <span class="n">xs</span><span class="p">):</span>
</span></span><span class="line"><span class="cl">    <span class="n">res</span> <span class="o">=</span> <span class="n">np</span><span class="o">.</span><span class="n">empty</span><span class="p">(</span><span class="nb">len</span><span class="p">(</span><span class="n">xs</span><span class="p">),</span> <span class="n">dtype</span><span class="o">=</span><span class="n">np</span><span class="o">.</span><span class="n">int32</span><span class="p">)</span>
</span></span><span class="line"><span class="cl">
</span></span><span class="line"><span class="cl">    <span class="k">for</span> <span class="n">i</span><span class="p">,</span> <span class="n">x</span> <span class="ow">in</span> <span class="nb">enumerate</span><span class="p">(</span><span class="n">xs</span><span class="p">):</span>
</span></span><span class="line"><span class="cl">        <span class="n">lo</span> <span class="o">=</span> <span class="mi">0</span>
</span></span><span class="line"><span class="cl">        <span class="n">hi</span> <span class="o">=</span> <span class="nb">len</span><span class="p">(</span><span class="n">a</span><span class="p">)</span>
</span></span><span class="line"><span class="cl">
</span></span><span class="line"><span class="cl">        <span class="k">while</span> <span class="n">lo</span> <span class="o">&lt;</span> <span class="n">hi</span><span class="p">:</span>
</span></span><span class="line"><span class="cl">            <span class="n">mid</span> <span class="o">=</span> <span class="p">(</span><span class="n">lo</span> <span class="o">+</span> <span class="n">hi</span><span class="p">)</span> <span class="o">//</span> <span class="mi">2</span>
</span></span><span class="line"><span class="cl">            <span class="k">if</span> <span class="n">a</span><span class="p">[</span><span class="n">mid</span><span class="p">]</span> <span class="o">&lt;</span> <span class="n">x</span><span class="p">:</span>
</span></span><span class="line"><span class="cl">                <span class="n">lo</span> <span class="o">=</span> <span class="n">mid</span> <span class="o">+</span> <span class="mi">1</span>
</span></span><span class="line"><span class="cl">            <span class="k">else</span><span class="p">:</span>
</span></span><span class="line"><span class="cl">                <span class="n">hi</span> <span class="o">=</span> <span class="n">mid</span>
</span></span><span class="line"><span class="cl">
</span></span><span class="line"><span class="cl">        <span class="n">res</span><span class="p">[</span><span class="n">i</span><span class="p">]</span> <span class="o">=</span> <span class="n">lo</span>
</span></span><span class="line"><span class="cl">
</span></span><span class="line"><span class="cl">    <span class="k">return</span> <span class="n">res</span></span></span></code></pre>
</div>
<p><img src="/numpy/searchsorted/images/figure1-fs8.png" alt=""></p>
<p>Running time per query grows logarithmically as the input size grows (note the logarithmic scale of x-axis).</p>
<p>For benchmarking, we generated two random arrays of uniformly distributed <code>np.int32</code> integers. The keys (i.e. the elements being searched) had a fixed length of 10,000, while the length of the sorted array varied up to $2^{30}$ ($4\ GiB$). Both keys and values arrays are contiguous in memory. Each benchmark was repeated 50 times, and we report the minimum execution time.</p>
<p>The benchmarks were run on a MacBook Pro with an Apple M1 Pro and 32 GB of memory. The M1 Pro has 128 KB of L1 data cache per performance core, enough to hold $2^{15}$ 32-bit integers, and a 12 MB L2 cache shared by its performance cores, enough to hold $1.5 * 2^{21}$ 32-bit integers.</p>
<h2 id="batching-with-numpy-arrays">Batching with NumPy arrays<a class="headerlink" href="#batching-with-numpy-arrays" title="Link to this heading">#</a></h2>
<p>The baseline implementation performs one binary search per query. Each search is independent, but executed sequentially in Python. We can adapt the algorithm so multiple binary searches make progress together in batches. For that we can represent the state of all searches as arrays and update them simultaneously using vectorized operations.</p>
<p>In NumPy, operations on arrays are executed in compiled C++ loops. This removes Python overhead and allows the CPU to efficiently process large batches of independent work.</p>
<h2 id="lets-vectorize-our-binary-search">Let&rsquo;s vectorize our binary search<a class="headerlink" href="#lets-vectorize-our-binary-search" title="Link to this heading">#</a></h2>
<p>To vectorize the algorithm, we reinterpret scalar variables as array state. Instead of a single <code>lo</code> and <code>hi</code>, we maintain one value per query. In our original implementation, the state consists of <code>lo</code>, <code>hi</code>, and <code>res</code>. Since <code>lo</code> ends up containing the final result, we can focus on tracking just <code>lo</code> and <code>hi</code>.</p>
<p>This first vectorized implementation is a direct translation of the previous algorithm.</p>


<div class="highlight">
  <pre class="chroma"><code><span class="line"><span class="cl"><span class="k">def</span> <span class="nf">searchsorted_py_np</span><span class="p">(</span><span class="n">a</span><span class="p">,</span> <span class="n">xs</span><span class="p">):</span>
</span></span><span class="line"><span class="cl">    <span class="n">lo</span> <span class="o">=</span> <span class="n">np</span><span class="o">.</span><span class="n">zeros</span><span class="p">(</span><span class="n">xs</span><span class="o">.</span><span class="n">shape</span><span class="p">,</span> <span class="n">dtype</span><span class="o">=</span><span class="n">np</span><span class="o">.</span><span class="n">int32</span><span class="p">)</span>
</span></span><span class="line"><span class="cl">    <span class="n">hi</span> <span class="o">=</span> <span class="n">np</span><span class="o">.</span><span class="n">full</span><span class="p">(</span><span class="n">xs</span><span class="o">.</span><span class="n">shape</span><span class="p">,</span> <span class="nb">len</span><span class="p">(</span><span class="n">a</span><span class="p">),</span> <span class="n">dtype</span><span class="o">=</span><span class="n">np</span><span class="o">.</span><span class="n">int32</span><span class="p">)</span>
</span></span><span class="line"><span class="cl">
</span></span><span class="line"><span class="cl">    <span class="k">while</span> <span class="kc">True</span><span class="p">:</span>
</span></span><span class="line"><span class="cl">        <span class="c1"># True for each element position where lo_i &lt; hi_i</span>
</span></span><span class="line"><span class="cl">        <span class="n">active</span> <span class="o">=</span> <span class="n">lo</span> <span class="o">&lt;</span> <span class="n">hi</span>
</span></span><span class="line"><span class="cl">        <span class="k">if</span> <span class="ow">not</span> <span class="n">np</span><span class="o">.</span><span class="n">any</span><span class="p">(</span><span class="n">active</span><span class="p">):</span>
</span></span><span class="line"><span class="cl">            <span class="c1"># this is basically `while lo &lt; hi:` in the pure-Python version</span>
</span></span><span class="line"><span class="cl">            <span class="k">break</span>
</span></span><span class="line"><span class="cl">
</span></span><span class="line"><span class="cl">        <span class="n">mid</span> <span class="o">=</span> <span class="p">(</span><span class="n">lo</span> <span class="o">+</span> <span class="n">hi</span><span class="p">)</span> <span class="o">//</span> <span class="mi">2</span>
</span></span><span class="line"><span class="cl">
</span></span><span class="line"><span class="cl">        <span class="c1"># only index those entries where active is true, so we don&#39;t modify already computed positions</span>
</span></span><span class="line"><span class="cl">        <span class="n">mid_a</span> <span class="o">=</span> <span class="n">mid</span><span class="p">[</span><span class="n">active</span><span class="p">]</span>
</span></span><span class="line"><span class="cl">        <span class="n">xs_a</span> <span class="o">=</span> <span class="n">xs</span><span class="p">[</span><span class="n">active</span><span class="p">]</span>
</span></span><span class="line"><span class="cl">
</span></span><span class="line"><span class="cl">        <span class="n">mask</span> <span class="o">=</span> <span class="n">a</span><span class="p">[</span><span class="n">mid_a</span><span class="p">]</span> <span class="o">&lt;</span> <span class="n">xs_a</span>
</span></span><span class="line"><span class="cl">
</span></span><span class="line"><span class="cl">        <span class="n">lo</span><span class="p">[</span><span class="n">active</span><span class="p">]</span> <span class="o">=</span> <span class="n">np</span><span class="o">.</span><span class="n">where</span><span class="p">(</span><span class="n">mask</span><span class="p">,</span> <span class="n">mid_a</span> <span class="o">+</span> <span class="mi">1</span><span class="p">,</span> <span class="n">lo</span><span class="p">[</span><span class="n">active</span><span class="p">])</span>
</span></span><span class="line"><span class="cl">        <span class="n">hi</span><span class="p">[</span><span class="n">active</span><span class="p">]</span> <span class="o">=</span> <span class="n">np</span><span class="o">.</span><span class="n">where</span><span class="p">(</span><span class="n">mask</span><span class="p">,</span> <span class="n">hi</span><span class="p">[</span><span class="n">active</span><span class="p">],</span> <span class="n">mid_a</span><span class="p">)</span>
</span></span><span class="line"><span class="cl">
</span></span><span class="line"><span class="cl">    <span class="k">return</span> <span class="n">lo</span></span></span></code></pre>
</div>
<p><img src="/numpy/searchsorted/images/figure2-fs8.png" alt=""></p>
<p>In this implementation, different queries can shrink their search intervals at different rates, so they may require different numbers of iterations to converge.</p>
<p>For example, consider searching for the keys <code>[-1, 2]</code> in the array <code>[0, 1]</code>. For the query <code>2</code>, the first iteration computes <code>mid = 1</code> and sets <code>lo = mid + 1 = 2</code>, so the interval becomes <code>[2, 2)</code> and the search converges in one step. For the query <code>-1</code>, the update instead sets <code>hi = mid = 1</code>, leaving the interval <code>[0, 1)</code> after the first step. This requires an additional iteration to collapse the interval to <code>[0, 0)</code>.</p>
<p>Because the searches can converge at different times, we need to keep track of which queries are still active. The active mask identifies the queries whose search intervals have not yet converged, allowing us to update only those queries.</p>
<h3 id="making-all-searches-take-the-same-number-of-steps">Making all searches take the same number of steps<a class="headerlink" href="#making-all-searches-take-the-same-number-of-steps" title="Link to this heading">#</a></h3>
<p>The important observation is that binary search does not actually need to terminate independently for each key. We can tweak each iteration update in a way that, once a search has converged, subsequent iterations can leave its interval unchanged.</p>
<p>We maintain the invariant that the insertion position lies in <code>[lo, hi)</code>. At each iteration, every interval is reduced to roughly half its previous size. After <code>np.ceil(np.log2(n))</code> iterations, every interval has collapsed to a single position.</p>
<p>For simplicity, this implementation assumes <code>a</code> is non-empty.</p>


<div class="highlight">
  <pre class="chroma"><code><span class="line"><span class="cl"><span class="k">def</span> <span class="nf">searchsorted_py_np_fixed</span><span class="p">(</span><span class="n">a</span><span class="p">,</span> <span class="n">xs</span><span class="p">):</span>
</span></span><span class="line"><span class="cl">    <span class="n">n</span> <span class="o">=</span> <span class="nb">len</span><span class="p">(</span><span class="n">a</span><span class="p">)</span>
</span></span><span class="line"><span class="cl">    <span class="n">lo</span> <span class="o">=</span> <span class="n">np</span><span class="o">.</span><span class="n">zeros</span><span class="p">(</span><span class="n">xs</span><span class="o">.</span><span class="n">shape</span><span class="p">,</span> <span class="n">dtype</span><span class="o">=</span><span class="n">np</span><span class="o">.</span><span class="n">int32</span><span class="p">)</span>
</span></span><span class="line"><span class="cl">    <span class="n">hi</span> <span class="o">=</span> <span class="n">np</span><span class="o">.</span><span class="n">full</span><span class="p">(</span><span class="n">xs</span><span class="o">.</span><span class="n">shape</span><span class="p">,</span> <span class="n">n</span><span class="p">,</span> <span class="n">dtype</span><span class="o">=</span><span class="n">np</span><span class="o">.</span><span class="n">int32</span><span class="p">)</span>
</span></span><span class="line"><span class="cl">
</span></span><span class="line"><span class="cl">    <span class="k">for</span> <span class="n">_</span> <span class="ow">in</span> <span class="nb">range</span><span class="p">(</span><span class="nb">int</span><span class="p">(</span><span class="n">np</span><span class="o">.</span><span class="n">ceil</span><span class="p">(</span><span class="n">np</span><span class="o">.</span><span class="n">log2</span><span class="p">(</span><span class="n">n</span><span class="p">)))):</span>
</span></span><span class="line"><span class="cl">        <span class="n">mid</span> <span class="o">=</span> <span class="p">(</span><span class="n">lo</span> <span class="o">+</span> <span class="n">hi</span><span class="p">)</span> <span class="o">//</span> <span class="mi">2</span>
</span></span><span class="line"><span class="cl">        <span class="n">go_left</span> <span class="o">=</span> <span class="n">xs</span> <span class="o">&lt;=</span> <span class="n">a</span><span class="p">[</span><span class="n">mid</span><span class="p">]</span>
</span></span><span class="line"><span class="cl">
</span></span><span class="line"><span class="cl">        <span class="c1"># for each position i:</span>
</span></span><span class="line"><span class="cl">        <span class="c1"># if go_left_i is True, we keep the `lo_i` value and `hi_i` is updated to `mid_i`</span>
</span></span><span class="line"><span class="cl">        <span class="c1"># if go_left_i is False, we keep the `hi_i` value and `lo_i` is updated to `mid_i`</span>
</span></span><span class="line"><span class="cl">        <span class="n">lo</span> <span class="o">=</span> <span class="n">np</span><span class="o">.</span><span class="n">where</span><span class="p">(</span><span class="n">go_left</span><span class="p">,</span> <span class="n">lo</span><span class="p">,</span> <span class="n">mid</span><span class="p">)</span>
</span></span><span class="line"><span class="cl">        <span class="n">hi</span> <span class="o">=</span> <span class="n">np</span><span class="o">.</span><span class="n">where</span><span class="p">(</span><span class="n">go_left</span><span class="p">,</span> <span class="n">mid</span><span class="p">,</span> <span class="n">hi</span><span class="p">)</span>
</span></span><span class="line"><span class="cl">
</span></span><span class="line"><span class="cl">    <span class="k">return</span> <span class="n">hi</span></span></span></code></pre>
</div>
<p>Removing the <code>active</code> tracker makes it up to 2× faster.</p>
<p><img src="/numpy/searchsorted/images/figure3-fs8.png" alt=""></p>
<p>A similar formulation can already be found in the Python ecosystem. For example, <a href="https://github.com/jax-ml/jax/blob/a6e4a8b95a731269bdf23e5b3e30da2f8494bb28/jax/_src/numpy/hijax.py#L330">JAX’s scan-based implementation</a></p>


<div class="highlight">
  <pre class="chroma"><code><span class="line"><span class="cl"><span class="o">...</span>
</span></span><span class="line"><span class="cl">
</span></span><span class="line"><span class="cl">
</span></span><span class="line"><span class="cl"><span class="k">def</span> <span class="nf">body_fun</span><span class="p">(</span><span class="n">state</span><span class="p">,</span> <span class="n">_</span><span class="p">):</span>
</span></span><span class="line"><span class="cl">    <span class="n">low</span><span class="p">,</span> <span class="n">high</span> <span class="o">=</span> <span class="n">state</span>
</span></span><span class="line"><span class="cl">    <span class="n">mid</span> <span class="o">=</span> <span class="n">low</span> <span class="o">+</span> <span class="p">(</span><span class="n">high</span> <span class="o">-</span> <span class="n">low</span><span class="p">)</span> <span class="o">//</span> <span class="mi">2</span>  <span class="c1"># use this form to avoid overflow</span>
</span></span><span class="line"><span class="cl">    <span class="n">go_left</span> <span class="o">=</span> <span class="n">op</span><span class="p">(</span><span class="n">query</span><span class="p">,</span> <span class="n">sorted_arr</span><span class="p">[</span><span class="n">mid</span><span class="p">])</span>
</span></span><span class="line"><span class="cl">    <span class="k">return</span> <span class="p">(</span><span class="n">lax</span><span class="o">.</span><span class="n">select</span><span class="p">(</span><span class="n">go_left</span><span class="p">,</span> <span class="n">low</span><span class="p">,</span> <span class="n">mid</span><span class="p">),</span> <span class="n">lax</span><span class="o">.</span><span class="n">select</span><span class="p">(</span><span class="n">go_left</span><span class="p">,</span> <span class="n">mid</span><span class="p">,</span> <span class="n">high</span><span class="p">)),</span> <span class="p">()</span>
</span></span><span class="line"><span class="cl">
</span></span><span class="line"><span class="cl">
</span></span><span class="line"><span class="cl"><span class="n">n_levels</span> <span class="o">=</span> <span class="nb">int</span><span class="p">(</span><span class="n">np</span><span class="o">.</span><span class="n">ceil</span><span class="p">(</span><span class="n">np</span><span class="o">.</span><span class="n">log2</span><span class="p">(</span><span class="n">n</span> <span class="o">+</span> <span class="mi">1</span><span class="p">)))</span>
</span></span><span class="line"><span class="cl"><span class="o">...</span></span></span></code></pre>
</div>
<h3 id="can-numpy-beat-numpy">Can NumPy beat NumPy?<a class="headerlink" href="#can-numpy-beat-numpy" title="Link to this heading">#</a></h3>
<p>Let’s compare the performance of this vectorized implementation with NumPy’s native <code>searchsorted</code> (using NumPy 2.4).</p>
<p><img src="/numpy/searchsorted/images/figure4-fs8.png" alt=""></p>
<p>Our vectorized Python implementation can be an order of magnitude faster than the native one for inputs with several keys. To understand why, let&rsquo;s take a look at <code>NumPy 2.4</code> implementation:</p>


<div class="highlight">
  <pre class="chroma"><code><span class="line"><span class="cl"><span class="k">template</span> <span class="o">&lt;</span><span class="k">class</span> <span class="nc">Tag</span><span class="p">,</span> <span class="n">side_t</span> <span class="n">side</span><span class="o">&gt;</span>
</span></span><span class="line"><span class="cl"><span class="k">static</span> <span class="kt">void</span>
</span></span><span class="line"><span class="cl"><span class="n">binsearch</span><span class="p">(</span><span class="k">const</span> <span class="kt">char</span> <span class="o">*</span><span class="n">arr</span><span class="p">,</span> <span class="k">const</span> <span class="kt">char</span> <span class="o">*</span><span class="n">key</span><span class="p">,</span> <span class="kt">char</span> <span class="o">*</span><span class="n">ret</span><span class="p">,</span> <span class="n">npy_intp</span> <span class="n">arr_len</span><span class="p">,</span>
</span></span><span class="line"><span class="cl">          <span class="n">npy_intp</span> <span class="n">key_len</span><span class="p">,</span> <span class="n">npy_intp</span> <span class="n">arr_str</span><span class="p">,</span> <span class="n">npy_intp</span> <span class="n">key_str</span><span class="p">,</span>
</span></span><span class="line"><span class="cl">          <span class="n">npy_intp</span> <span class="n">ret_str</span><span class="p">,</span> <span class="n">PyArrayObject</span> <span class="o">*</span><span class="p">)</span>
</span></span><span class="line"><span class="cl"><span class="p">{</span>
</span></span><span class="line"><span class="cl">    <span class="k">using</span> <span class="n">T</span> <span class="o">=</span> <span class="k">typename</span> <span class="n">Tag</span><span class="o">::</span><span class="n">type</span><span class="p">;</span>
</span></span><span class="line"><span class="cl">    <span class="k">auto</span> <span class="n">cmp</span> <span class="o">=</span> <span class="n">side_to_cmp</span><span class="o">&lt;</span><span class="n">Tag</span><span class="p">,</span> <span class="n">side</span><span class="o">&gt;::</span><span class="n">value</span><span class="p">;</span>
</span></span><span class="line"><span class="cl">    <span class="n">npy_intp</span> <span class="n">min_idx</span> <span class="o">=</span> <span class="mi">0</span><span class="p">;</span>
</span></span><span class="line"><span class="cl">    <span class="n">npy_intp</span> <span class="n">max_idx</span> <span class="o">=</span> <span class="n">arr_len</span><span class="p">;</span>
</span></span><span class="line"><span class="cl">    <span class="n">T</span> <span class="n">last_key_val</span><span class="p">;</span>
</span></span><span class="line"><span class="cl">
</span></span><span class="line"><span class="cl">    <span class="k">if</span> <span class="p">(</span><span class="n">key_len</span> <span class="o">==</span> <span class="mi">0</span><span class="p">)</span> <span class="p">{</span>
</span></span><span class="line"><span class="cl">        <span class="k">return</span><span class="p">;</span>
</span></span><span class="line"><span class="cl">    <span class="p">}</span>
</span></span><span class="line"><span class="cl">    <span class="n">last_key_val</span> <span class="o">=</span> <span class="o">*</span><span class="p">(</span><span class="k">const</span> <span class="n">T</span> <span class="o">*</span><span class="p">)</span><span class="n">key</span><span class="p">;</span>
</span></span><span class="line"><span class="cl">
</span></span><span class="line"><span class="cl">    <span class="k">for</span> <span class="p">(;</span> <span class="n">key_len</span> <span class="o">&gt;</span> <span class="mi">0</span><span class="p">;</span> <span class="n">key_len</span><span class="o">--</span><span class="p">,</span> <span class="n">key</span> <span class="o">+=</span> <span class="n">key_str</span><span class="p">,</span> <span class="n">ret</span> <span class="o">+=</span> <span class="n">ret_str</span><span class="p">)</span> <span class="p">{</span>
</span></span><span class="line"><span class="cl">        <span class="k">const</span> <span class="n">T</span> <span class="n">key_val</span> <span class="o">=</span> <span class="o">*</span><span class="p">(</span><span class="k">const</span> <span class="n">T</span> <span class="o">*</span><span class="p">)</span><span class="n">key</span><span class="p">;</span>
</span></span><span class="line"><span class="cl">        <span class="cm">/*
</span></span></span><span class="line"><span class="cl"><span class="cm">         * Updating only one of the indices based on the previous key
</span></span></span><span class="line"><span class="cl"><span class="cm">         * gives the search a big boost when keys are sorted, but slightly
</span></span></span><span class="line"><span class="cl"><span class="cm">         * slows down things for purely random ones.
</span></span></span><span class="line"><span class="cl"><span class="cm">         */</span>
</span></span><span class="line"><span class="cl">        <span class="k">if</span> <span class="p">(</span><span class="n">cmp</span><span class="p">(</span><span class="n">last_key_val</span><span class="p">,</span> <span class="n">key_val</span><span class="p">))</span> <span class="p">{</span>
</span></span><span class="line"><span class="cl">            <span class="n">max_idx</span> <span class="o">=</span> <span class="n">arr_len</span><span class="p">;</span>
</span></span><span class="line"><span class="cl">        <span class="p">}</span>
</span></span><span class="line"><span class="cl">        <span class="k">else</span> <span class="p">{</span>
</span></span><span class="line"><span class="cl">            <span class="n">min_idx</span> <span class="o">=</span> <span class="mi">0</span><span class="p">;</span>
</span></span><span class="line"><span class="cl">            <span class="n">max_idx</span> <span class="o">=</span> <span class="p">(</span><span class="n">max_idx</span> <span class="o">&lt;</span> <span class="n">arr_len</span><span class="p">)</span> <span class="o">?</span> <span class="p">(</span><span class="n">max_idx</span> <span class="o">+</span> <span class="mi">1</span><span class="p">)</span> <span class="o">:</span> <span class="n">arr_len</span><span class="p">;</span>
</span></span><span class="line"><span class="cl">        <span class="p">}</span>
</span></span><span class="line"><span class="cl">
</span></span><span class="line"><span class="cl">        <span class="n">last_key_val</span> <span class="o">=</span> <span class="n">key_val</span><span class="p">;</span>
</span></span><span class="line"><span class="cl">
</span></span><span class="line"><span class="cl">        <span class="k">while</span> <span class="p">(</span><span class="n">min_idx</span> <span class="o">&lt;</span> <span class="n">max_idx</span><span class="p">)</span> <span class="p">{</span>
</span></span><span class="line"><span class="cl">            <span class="k">const</span> <span class="n">npy_intp</span> <span class="n">mid_idx</span> <span class="o">=</span> <span class="n">min_idx</span> <span class="o">+</span> <span class="p">((</span><span class="n">max_idx</span> <span class="o">-</span> <span class="n">min_idx</span><span class="p">)</span> <span class="o">&gt;&gt;</span> <span class="mi">1</span><span class="p">);</span>
</span></span><span class="line"><span class="cl">            <span class="k">const</span> <span class="n">T</span> <span class="n">mid_val</span> <span class="o">=</span> <span class="o">*</span><span class="p">(</span><span class="k">const</span> <span class="n">T</span> <span class="o">*</span><span class="p">)(</span><span class="n">arr</span> <span class="o">+</span> <span class="n">mid_idx</span> <span class="o">*</span> <span class="n">arr_str</span><span class="p">);</span>
</span></span><span class="line"><span class="cl">            <span class="k">if</span> <span class="p">(</span><span class="n">cmp</span><span class="p">(</span><span class="n">mid_val</span><span class="p">,</span> <span class="n">key_val</span><span class="p">))</span> <span class="p">{</span>
</span></span><span class="line"><span class="cl">                <span class="n">min_idx</span> <span class="o">=</span> <span class="n">mid_idx</span> <span class="o">+</span> <span class="mi">1</span><span class="p">;</span>
</span></span><span class="line"><span class="cl">            <span class="p">}</span>
</span></span><span class="line"><span class="cl">            <span class="k">else</span> <span class="p">{</span>
</span></span><span class="line"><span class="cl">                <span class="n">max_idx</span> <span class="o">=</span> <span class="n">mid_idx</span><span class="p">;</span>
</span></span><span class="line"><span class="cl">            <span class="p">}</span>
</span></span><span class="line"><span class="cl">        <span class="p">}</span>
</span></span><span class="line"><span class="cl">        <span class="o">*</span><span class="p">(</span><span class="n">npy_intp</span> <span class="o">*</span><span class="p">)</span><span class="n">ret</span> <span class="o">=</span> <span class="n">min_idx</span><span class="p">;</span>
</span></span><span class="line"><span class="cl">    <span class="p">}</span>
</span></span><span class="line"><span class="cl"><span class="p">}</span></span></span></code></pre>
</div>
<p>Ignoring the pointer arithmetic details, the core algorithm is a classic binary search executed independently for each key. There is also an optimization that reuses previous search bounds when the input keys are sorted.</p>
<p>This implementation performs one binary search per key, where each search is a fully sequential process. Each iteration of the binary search depends on the result of the previous one (the midpoint determines which part of the array is inspected next). This creates a dependency chain within each search: the next memory access depends on the result of the previous comparison.</p>
<p>For large arrays, binary-search reads also tend to be cache-unfriendly, since each step may require a read from a different cache line. Cache misses have a greater impact on the sequential implementation because each step may stall waiting for the previous memory access to complete.</p>
<p>The vectorized implementation performs the same logical step across all queries at once (all queries advance at each step together). <strong>With multiple independent searches, the CPU can have several memory accesses in flight at once.</strong> This aligns with the observed running time once the array size exceeds the L1 and L2 cache sizes.</p>
<h3 id="can-we-optimize-numpy">Can we optimize NumPy?<a class="headerlink" href="#can-we-optimize-numpy" title="Link to this heading">#</a></h3>
<p>The previous vectorized implementation maintains two arrays, <code>lo</code> and <code>hi</code>, to represent the search interval for each query. If we were to port this exact implementation into NumPy natively, it would require using $O(K)$ additional memory where <code>K</code> is the number of queries. Even though this approach is potentially faster, this is unacceptable for memory-sensitive workloads.</p>
<p>To reduce the state required, we can reformulate binary search in terms of interval boundaries. Instead of tracking both <code>lo</code> and <code>hi</code> for each query, we describe each interval using its left boundary and its length.</p>
<p>The key observation is that if we structure the algorithm so that all queries shrink their intervals by the same amount at each iteration, then every interval has the same length at each iteration. This means we do not need to store a separate <code>hi</code> per query: it can be reconstructed from a single array <code>lo</code> and a global length. Note that this still needs $O(K)$ space for the output, but it requires only $O(1)$ memory beyond that output.</p>
<p>This gives us the following Python implementation:</p>


<div class="highlight">
  <pre class="chroma"><code><span class="line"><span class="cl"><span class="k">def</span> <span class="nf">searchsorted_py_np_fast_where</span><span class="p">(</span><span class="n">arr</span><span class="p">,</span> <span class="n">keys</span><span class="p">):</span>
</span></span><span class="line"><span class="cl">    <span class="n">K</span> <span class="o">=</span> <span class="n">keys</span><span class="o">.</span><span class="n">shape</span><span class="p">[</span><span class="mi">0</span><span class="p">]</span>
</span></span><span class="line"><span class="cl">    <span class="n">length</span> <span class="o">=</span> <span class="n">arr</span><span class="o">.</span><span class="n">shape</span><span class="p">[</span><span class="mi">0</span><span class="p">]</span>
</span></span><span class="line"><span class="cl">
</span></span><span class="line"><span class="cl">    <span class="n">base</span> <span class="o">=</span> <span class="n">np</span><span class="o">.</span><span class="n">zeros</span><span class="p">(</span><span class="n">K</span><span class="p">,</span> <span class="n">dtype</span><span class="o">=</span><span class="n">np</span><span class="o">.</span><span class="n">intp</span><span class="p">)</span>
</span></span><span class="line"><span class="cl">
</span></span><span class="line"><span class="cl">    <span class="c1"># Invariant: the insertion index lies in [base, base + length]</span>
</span></span><span class="line"><span class="cl">    <span class="k">while</span> <span class="n">length</span> <span class="o">&gt;</span> <span class="mi">1</span><span class="p">:</span>
</span></span><span class="line"><span class="cl">        <span class="n">half</span> <span class="o">=</span> <span class="n">length</span> <span class="o">&gt;&gt;</span> <span class="mi">1</span>
</span></span><span class="line"><span class="cl">        <span class="n">mid</span> <span class="o">=</span> <span class="n">base</span> <span class="o">+</span> <span class="n">half</span>
</span></span><span class="line"><span class="cl">        <span class="n">base</span> <span class="o">=</span> <span class="n">np</span><span class="o">.</span><span class="n">where</span><span class="p">(</span><span class="n">keys</span> <span class="o">&gt;</span> <span class="n">arr</span><span class="p">[</span><span class="n">mid</span><span class="p">],</span> <span class="n">mid</span><span class="p">,</span> <span class="n">base</span><span class="p">)</span>
</span></span><span class="line"><span class="cl">        <span class="n">length</span> <span class="o">-=</span> <span class="n">half</span>
</span></span><span class="line"><span class="cl">
</span></span><span class="line"><span class="cl">    <span class="c1"># Final step: result is either base and base + 1</span>
</span></span><span class="line"><span class="cl">    <span class="n">base</span> <span class="o">=</span> <span class="n">np</span><span class="o">.</span><span class="n">where</span><span class="p">(</span><span class="n">keys</span> <span class="o">&gt;</span> <span class="n">arr</span><span class="p">[</span><span class="n">base</span><span class="p">],</span> <span class="n">base</span> <span class="o">+</span> <span class="mi">1</span><span class="p">,</span> <span class="n">base</span><span class="p">)</span>
</span></span><span class="line"><span class="cl">
</span></span><span class="line"><span class="cl">    <span class="k">return</span> <span class="n">base</span></span></span></code></pre>
</div>
<p>Only a single array <code>base</code> is needed to store the per-query state, while length is shared across all queries. The output is written directly into base, so no additional state array is required. This implementation still requires allocating a temporary <code>mid</code> to hold the midpoints, but this can be avoided in the C++ port.</p>
<p>This formulation is closely related to the branchless binary search approach discussed in <a href="https://en.algorithmica.org/hpc/data-structures/binary-search/#removing-branches">Algorithmica&rsquo;s case study</a>. In the formulation we use, the invariant range is <code>[base, base + length]</code>. Therefore a final step is required when <code>length = 1</code> to resolve whether the insertion point falls to the left or right of <code>base</code>.</p>
<p>The reformulated implementation is significantly faster:</p>
<p><img src="/numpy/searchsorted/images/figure6-fs8.png" alt=""></p>
<h3 id="porting-it-into-c">Porting it into C++<a class="headerlink" href="#porting-it-into-c" title="Link to this heading">#</a></h3>
<p>The performance results show the benefit of reducing the per-query state. We can now translate it almost directly to C++ with $O(1)$ additional memory.</p>


<div class="highlight">
  <pre class="chroma"><code><span class="line"><span class="cl"><span class="k">template</span> <span class="o">&lt;</span><span class="k">class</span> <span class="nc">Tag</span><span class="p">,</span> <span class="n">side_t</span> <span class="n">side</span><span class="o">&gt;</span>
</span></span><span class="line"><span class="cl"><span class="k">static</span> <span class="kt">void</span>
</span></span><span class="line"><span class="cl"><span class="n">binsearch</span><span class="p">(</span><span class="k">const</span> <span class="kt">char</span> <span class="o">*</span><span class="n">arr</span><span class="p">,</span> <span class="k">const</span> <span class="kt">char</span> <span class="o">*</span><span class="n">key</span><span class="p">,</span> <span class="kt">char</span> <span class="o">*</span><span class="n">ret</span><span class="p">,</span> <span class="n">npy_intp</span> <span class="n">arr_len</span><span class="p">,</span>
</span></span><span class="line"><span class="cl">          <span class="n">npy_intp</span> <span class="n">key_len</span><span class="p">,</span> <span class="n">npy_intp</span> <span class="n">arr_str</span><span class="p">,</span> <span class="n">npy_intp</span> <span class="n">key_str</span><span class="p">,</span>
</span></span><span class="line"><span class="cl">          <span class="n">npy_intp</span> <span class="n">ret_str</span><span class="p">,</span> <span class="n">PyArrayObject</span> <span class="o">*</span><span class="p">)</span>
</span></span><span class="line"><span class="cl"><span class="p">{</span>
</span></span><span class="line"><span class="cl">    <span class="k">using</span> <span class="n">T</span> <span class="o">=</span> <span class="k">typename</span> <span class="n">Tag</span><span class="o">::</span><span class="n">type</span><span class="p">;</span>
</span></span><span class="line"><span class="cl">    <span class="k">auto</span> <span class="n">cmp</span> <span class="o">=</span> <span class="n">side_to_cmp</span><span class="o">&lt;</span><span class="n">Tag</span><span class="p">,</span> <span class="n">side</span><span class="o">&gt;::</span><span class="n">value</span><span class="p">;</span>
</span></span><span class="line"><span class="cl">
</span></span><span class="line"><span class="cl">    <span class="c1">// If the array length is 0 we return all 0s
</span></span></span><span class="line"><span class="cl"><span class="c1"></span>    <span class="k">if</span> <span class="p">(</span><span class="n">arr_len</span> <span class="o">&lt;=</span> <span class="mi">0</span><span class="p">)</span> <span class="p">{</span>
</span></span><span class="line"><span class="cl">        <span class="k">for</span> <span class="p">(</span><span class="n">npy_intp</span> <span class="n">i</span> <span class="o">=</span> <span class="mi">0</span><span class="p">;</span> <span class="n">i</span> <span class="o">&lt;</span> <span class="n">key_len</span><span class="p">;</span> <span class="o">++</span><span class="n">i</span><span class="p">)</span> <span class="p">{</span>
</span></span><span class="line"><span class="cl">            <span class="o">*</span><span class="p">(</span><span class="n">npy_intp</span> <span class="o">*</span><span class="p">)(</span><span class="n">ret</span> <span class="o">+</span> <span class="n">i</span> <span class="o">*</span> <span class="n">ret_str</span><span class="p">)</span> <span class="o">=</span> <span class="mi">0</span><span class="p">;</span>
</span></span><span class="line"><span class="cl">        <span class="p">}</span>
</span></span><span class="line"><span class="cl">        <span class="k">return</span><span class="p">;</span>
</span></span><span class="line"><span class="cl">    <span class="p">}</span>
</span></span><span class="line"><span class="cl">
</span></span><span class="line"><span class="cl">    <span class="cm">/*
</span></span></span><span class="line"><span class="cl"><span class="cm">    base = np.zeros(K, dtype=np.intp)
</span></span></span><span class="line"><span class="cl"><span class="cm">
</span></span></span><span class="line"><span class="cl"><span class="cm">    We unroll the first iteration for the following reasons:
</span></span></span><span class="line"><span class="cl"><span class="cm">        1. ret is not initialized with the bases, so we save |keys| writes
</span></span></span><span class="line"><span class="cl"><span class="cm">        by not having to initialize it with 0s.
</span></span></span><span class="line"><span class="cl"><span class="cm">        2. By assuming the initial base for every key is 0, we also save
</span></span></span><span class="line"><span class="cl"><span class="cm">        |keys| reads.
</span></span></span><span class="line"><span class="cl"><span class="cm">        3. In the first iteration, all elements are compared against the
</span></span></span><span class="line"><span class="cl"><span class="cm">        median. So we can store it in a variable and use it for all keys.
</span></span></span><span class="line"><span class="cl"><span class="cm">
</span></span></span><span class="line"><span class="cl"><span class="cm">    This initial block replaces the initialization loop that is used for the
</span></span></span><span class="line"><span class="cl"><span class="cm">    arr_len==0 case. Note that when arr_len = 1, then half is 0 so the
</span></span></span><span class="line"><span class="cl"><span class="cm">    following block initializes the array as with 0s.
</span></span></span><span class="line"><span class="cl"><span class="cm">    */</span>
</span></span><span class="line"><span class="cl">    <span class="n">npy_intp</span> <span class="n">interval_length</span> <span class="o">=</span> <span class="n">arr_len</span><span class="p">;</span>
</span></span><span class="line"><span class="cl">    <span class="n">npy_intp</span> <span class="n">half</span> <span class="o">=</span> <span class="n">interval_length</span> <span class="o">&gt;&gt;</span> <span class="mi">1</span><span class="p">;</span>
</span></span><span class="line"><span class="cl">    <span class="n">interval_length</span> <span class="o">-=</span> <span class="n">half</span><span class="p">;</span> <span class="c1">// length -&gt; ceil(length / 2)
</span></span></span><span class="line"><span class="cl"><span class="c1"></span>
</span></span><span class="line"><span class="cl">    <span class="n">npy_intp</span> <span class="n">base</span> <span class="o">=</span> <span class="mi">0</span><span class="p">;</span>
</span></span><span class="line"><span class="cl">    <span class="k">const</span> <span class="n">T</span> <span class="n">mid_val</span> <span class="o">=</span> <span class="o">*</span><span class="p">(</span><span class="k">const</span> <span class="n">T</span> <span class="o">*</span><span class="p">)(</span><span class="n">arr</span> <span class="o">+</span> <span class="p">(</span><span class="n">base</span> <span class="o">+</span> <span class="n">half</span><span class="p">)</span> <span class="o">*</span> <span class="n">arr_str</span><span class="p">);</span>
</span></span><span class="line"><span class="cl">
</span></span><span class="line"><span class="cl">    <span class="k">for</span> <span class="p">(</span><span class="n">npy_intp</span> <span class="n">i</span> <span class="o">=</span> <span class="mi">0</span><span class="p">;</span> <span class="n">i</span> <span class="o">&lt;</span> <span class="n">key_len</span><span class="p">;</span> <span class="o">++</span><span class="n">i</span><span class="p">)</span> <span class="p">{</span>
</span></span><span class="line"><span class="cl">        <span class="k">const</span> <span class="n">T</span> <span class="n">key_val</span> <span class="o">=</span> <span class="o">*</span><span class="p">(</span><span class="k">const</span> <span class="n">T</span> <span class="o">*</span><span class="p">)(</span><span class="n">key</span> <span class="o">+</span> <span class="n">i</span> <span class="o">*</span> <span class="n">key_str</span><span class="p">);</span>
</span></span><span class="line"><span class="cl">        <span class="o">*</span><span class="p">(</span><span class="n">npy_intp</span> <span class="o">*</span><span class="p">)(</span><span class="n">ret</span> <span class="o">+</span> <span class="n">i</span> <span class="o">*</span> <span class="n">ret_str</span><span class="p">)</span> <span class="o">=</span> <span class="n">cmp</span><span class="p">(</span><span class="n">mid_val</span><span class="p">,</span> <span class="n">key_val</span><span class="p">)</span> <span class="o">*</span> <span class="n">half</span><span class="p">;</span>
</span></span><span class="line"><span class="cl">    <span class="p">}</span>
</span></span><span class="line"><span class="cl">
</span></span><span class="line"><span class="cl">    <span class="cm">/*
</span></span></span><span class="line"><span class="cl"><span class="cm">        while length &gt; 1:
</span></span></span><span class="line"><span class="cl"><span class="cm">            half = length &gt;&gt; 1
</span></span></span><span class="line"><span class="cl"><span class="cm">            length -= half
</span></span></span><span class="line"><span class="cl"><span class="cm">            mid = base + half
</span></span></span><span class="line"><span class="cl"><span class="cm">            base = np.where(keys &gt; arr[mid], mid, base)
</span></span></span><span class="line"><span class="cl"><span class="cm">    */</span>
</span></span><span class="line"><span class="cl">    <span class="k">while</span> <span class="p">(</span><span class="n">interval_length</span> <span class="o">&gt;</span> <span class="mi">1</span><span class="p">)</span> <span class="p">{</span>
</span></span><span class="line"><span class="cl">        <span class="n">npy_intp</span> <span class="n">half</span> <span class="o">=</span> <span class="n">interval_length</span> <span class="o">&gt;&gt;</span> <span class="mi">1</span><span class="p">;</span>
</span></span><span class="line"><span class="cl">        <span class="n">interval_length</span> <span class="o">-=</span> <span class="n">half</span><span class="p">;</span>
</span></span><span class="line"><span class="cl">
</span></span><span class="line"><span class="cl">        <span class="k">for</span> <span class="p">(</span><span class="n">npy_intp</span> <span class="n">i</span> <span class="o">=</span> <span class="mi">0</span><span class="p">;</span> <span class="n">i</span> <span class="o">&lt;</span> <span class="n">key_len</span><span class="p">;</span> <span class="o">++</span><span class="n">i</span><span class="p">)</span> <span class="p">{</span>
</span></span><span class="line"><span class="cl">            <span class="n">npy_intp</span> <span class="o">&amp;</span><span class="n">base</span> <span class="o">=</span> <span class="o">*</span><span class="p">(</span><span class="n">npy_intp</span> <span class="o">*</span><span class="p">)(</span><span class="n">ret</span> <span class="o">+</span> <span class="n">i</span> <span class="o">*</span> <span class="n">ret_str</span><span class="p">);</span>
</span></span><span class="line"><span class="cl">            <span class="k">const</span> <span class="n">T</span> <span class="n">mid_val</span> <span class="o">=</span> <span class="o">*</span><span class="p">(</span><span class="k">const</span> <span class="n">T</span> <span class="o">*</span><span class="p">)(</span><span class="n">arr</span> <span class="o">+</span> <span class="p">(</span><span class="n">base</span> <span class="o">+</span> <span class="n">half</span><span class="p">)</span> <span class="o">*</span> <span class="n">arr_str</span><span class="p">);</span>
</span></span><span class="line"><span class="cl">            <span class="k">const</span> <span class="n">T</span> <span class="n">key_val</span> <span class="o">=</span> <span class="o">*</span><span class="p">(</span><span class="k">const</span> <span class="n">T</span> <span class="o">*</span><span class="p">)(</span><span class="n">key</span> <span class="o">+</span> <span class="n">i</span> <span class="o">*</span> <span class="n">key_str</span><span class="p">);</span>
</span></span><span class="line"><span class="cl">            <span class="n">base</span> <span class="o">+=</span> <span class="n">cmp</span><span class="p">(</span><span class="n">mid_val</span><span class="p">,</span> <span class="n">key_val</span><span class="p">)</span> <span class="o">*</span> <span class="n">half</span><span class="p">;</span>
</span></span><span class="line"><span class="cl">        <span class="p">}</span>
</span></span><span class="line"><span class="cl">    <span class="p">}</span>
</span></span><span class="line"><span class="cl">
</span></span><span class="line"><span class="cl">    <span class="c1">// base = np.where(keys &gt; arr[base], base + 1, base)
</span></span></span><span class="line"><span class="cl"><span class="c1"></span>    <span class="k">for</span> <span class="p">(</span><span class="n">npy_intp</span> <span class="n">i</span> <span class="o">=</span> <span class="mi">0</span><span class="p">;</span> <span class="n">i</span> <span class="o">&lt;</span> <span class="n">key_len</span><span class="p">;</span> <span class="o">++</span><span class="n">i</span><span class="p">)</span> <span class="p">{</span>
</span></span><span class="line"><span class="cl">        <span class="n">npy_intp</span> <span class="o">&amp;</span><span class="n">base</span> <span class="o">=</span> <span class="o">*</span><span class="p">(</span><span class="n">npy_intp</span> <span class="o">*</span><span class="p">)(</span><span class="n">ret</span> <span class="o">+</span> <span class="n">i</span> <span class="o">*</span> <span class="n">ret_str</span><span class="p">);</span>
</span></span><span class="line"><span class="cl">        <span class="k">const</span> <span class="n">T</span> <span class="n">key_val</span> <span class="o">=</span> <span class="o">*</span><span class="p">(</span><span class="k">const</span> <span class="n">T</span> <span class="o">*</span><span class="p">)(</span><span class="n">key</span> <span class="o">+</span> <span class="n">i</span> <span class="o">*</span> <span class="n">key_str</span><span class="p">);</span>
</span></span><span class="line"><span class="cl">        <span class="n">base</span> <span class="o">+=</span> <span class="n">cmp</span><span class="p">(</span><span class="o">*</span><span class="p">(</span><span class="k">const</span> <span class="n">T</span> <span class="o">*</span><span class="p">)(</span><span class="n">arr</span> <span class="o">+</span> <span class="n">base</span> <span class="o">*</span> <span class="n">arr_str</span><span class="p">),</span> <span class="n">key_val</span><span class="p">);</span>
</span></span><span class="line"><span class="cl">    <span class="p">}</span>
</span></span><span class="line"><span class="cl"><span class="p">}</span></span></span></code></pre>
</div>
<p>Note that we exploited a property of the first iteration of the binary search. Because the initial value of every result entry is implicitly zero, we can skip writing and reading those values during the first iteration. Moreover, in the first iteration all elements are compared against the same median, so we can read its value once instead of <code>K</code> times.</p>
<p>This implementation was ported directly into NumPy as part of PR <a href="https://github.com/numpy/numpy/pull/30517">#30517</a>, which was included in the <a href="https://numpy.org/devdocs/release/2.5.0-notes.html#improved-performance-of-numpy-searchsorted">2.5 release</a>. Now let&rsquo;s do a final comparison between NumPy 2.4 and 2.5, and our vectorized Python implementation:</p>
<p><img src="/numpy/searchsorted/images/figure7-fs8.png" alt=""></p>
<p>The native 2.5 version is slightly faster than the vectorized Python one and up to 25× faster than NumPy 2.4&rsquo;s implementation in our benchmarks.</p>
<h3 id="ecosystem-comparison">Ecosystem Comparison<a class="headerlink" href="#ecosystem-comparison" title="Link to this heading">#</a></h3>
<p>We can compare our optimized NumPy 2.5 against other libraries in the ecosystem. For this experiment, we selected the Python libraries JAX, TensorFlow, and PyTorch.</p>
<p>TensorFlow and PyTorch follow a different approach from JAX and NumPy. While JAX and NumPy leverage vectorized/batched operations to hide memory latency, TensorFlow and PyTorch parallelize independent searches across CPU threads. Search keys are partitioned into batches that are processed by different threads. For more details, see the <a href="https://github.com/pytorch/pytorch/blob/b1bb860d3c812371b89a9725407230216e7369b5/aten/src/ATen/native/Bucketization.cpp#L88">PyTorch</a> and <a href="https://github.com/tensorflow/tensorflow/blob/bb8d3f2443d70ec8c2aae1288fbf5782c771aa60/tensorflow/core/kernels/searchsorted_op.cc#L67">TensorFlow</a> implementations.</p>
<p>In the benchmarks, we limited parallelism to 8 cores and we increased the number of query keys from 10,000 to 20,000. This gives the multithreaded implementations enough independent work to amortize thread-scheduling overhead.</p>
<p><img src="/numpy/searchsorted/images/figure8-fs8.png" alt=""></p>
<p>The benchmark shows that NumPy is competitive with the selected libraries in our benchmarks. All implementations exhibit similar behavior once the search array grows beyond the CPU cache.</p>
<p>If we disable multithreading, the performance of PyTorch and TensorFlow degrades, and both exhibit a similar trend to NumPy 2.4&rsquo;s implementation. Once the search array grows beyond the CPU cache, the cost of memory accesses dominates.</p>
<p><img src="/numpy/searchsorted/images/figure9-fs8.png" alt=""></p>
<p>It would be worth benchmarking whether both techniques could be combined: batching binary searches within each thread. However, once the memory subsystem becomes saturated, additional cores can compete for the same memory bandwidth. At that point, improving the memory access patterns may be a more promising direction, for example by using a different layout such as the Eytzinger layout (discussed in detail in the <a href="https://en.algorithmica.org/hpc/data-structures/binary-search/#eytzinger-layout">Algorithmica book</a>).</p>
<h3 id="conclusion">Conclusion<a class="headerlink" href="#conclusion" title="Link to this heading">#</a></h3>
<p>We made <code>np.searchsorted</code> up to 25x faster in our benchmarks. Given NumPy&rsquo;s reach in the Python ecosystem, this optimization will benefit several libraries that depend on it. Other libraries in the Python ecosystem with their own binary search implementation may also benefit from adopting similar batching techniques.</p>
<p>Interestingly, we used NumPy array primitives to derive an initial Python implementation that outperformed NumPy 2.4&rsquo;s implementation. This shows how powerful NumPy&rsquo;s array primitives can be for implementing highly performant algorithms. <strong>A vectorized NumPy implementation in Python can outperform a scalar native implementation by exploiting independent work</strong>.</p>
<p>Cache-friendly layouts such as the Eytzinger layout are another interesting direction for making <code>searchsorted</code> faster. It would be interesting to explore whether the array API could expose such layouts through an interface like <code>searchsorted(arr, keys, layout=&quot;eytzinger&quot;)</code>, although this would require carefully defining the API semantics since the Eytzinger representation is not sorted.</p>
]]></content>
            
                 
                    
                 
                    
                         
                        
                            
                             
                                <category scheme="taxonomy:Tags" term="numpy" label="numpy" />
                             
                                <category scheme="taxonomy:Tags" term="performance" label="performance" />
                             
                                <category scheme="taxonomy:Tags" term="binary-search" label="binary-search" />
                             
                                <category scheme="taxonomy:Tags" term="searchsorted" label="searchsorted" />
                            
                        
                    
                
            
        </entry>
    
</feed>
