Can I perform binary tree search with the standard library without wrapping the float type and abusing the BTreeMap?

Viewed 2073

I would like to find the first element which is greater than a limit from an ordered collection. While iteration over it is always an option, I need a faster one. Currently, I came up with a solution like this but it feels a little hacky:

use std::cmp::Ordering;
use std::collections::BTreeMap;
use std::ops::Bound::{Included, Unbounded};

#[derive(Debug)]
struct FloatWrapper(f32);

impl Eq for FloatWrapper {}

impl PartialEq for FloatWrapper {
    fn eq(&self, other: &Self) -> bool {
        (self.0 - other.0).abs() < 1.17549435e-36f32
    }
}

impl Ord for FloatWrapper {
    fn cmp(&self, other: &Self) -> Ordering {
        if (self.0 - other.0).abs() < 1.17549435e-36f32 {
            Ordering::Equal
        } else if self.0 - other.0 > 0.0 {
            Ordering::Greater
        } else if self.0 - other.0 < 0.0 {
            Ordering::Less
        } else {
            Ordering::Equal
        }
    }
}

impl PartialOrd for FloatWrapper {
    fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
        Some(self.cmp(other))
    }
}
  • The wrapper around the float is not nice even that I am sure that there will be no NaNs

  • The Range is also unnecessary since I want a single element.

Is there a better way of achieving a similar result using only Rust's standard library? I know that there are plenty of tree implementations but it feels like overkill.

After the suggestions in the answer to use the iterator I did a little benchmark with the following code:

fn main() {
    let measure = vec![
        10, 15, 20, 30, 40, 50, 60, 70, 80, 90, 100, 110, 120, 130, 140, 150, 160, 170, 180, 190,
        200,
    ];

    let mut measured_binary = Vec::new();
    let mut measured_iter = Vec::new();
    let mut measured_vec = Vec::new();

    for size in measure {
        let mut ww = BTreeMap::new();
        let mut what_found = Vec::new();
        for _ in 0..size {
            let now: f32 = thread_rng().gen_range(0.0, 1.0);
            ww.insert(FloatWrapper(now), now);
        }
        let what_to_search: Vec<FloatWrapper> = (0..10000)
            .map(|_| thread_rng().gen_range(0.0, 0.8))
            .map(|x| FloatWrapper(x))
            .collect();
        let mut rez = 0;

        for current in &what_to_search {
            let now = Instant::now();
            let m = find_one(&ww, current);
            rez += now.elapsed().as_nanos();
            what_found.push(m);
        }

        measured_binary.push(rez);
        rez = 0;

        for current in &what_to_search {
            let now = Instant::now();
            let m = find_two(&ww, current);
            rez += now.elapsed().as_nanos();
            what_found.push(m);
        }
        measured_iter.push(rez);

        let ww_in_vec: Vec<(FloatWrapper, f32)> =
            ww.iter().map(|(&key, &value)| (key, value)).collect();

        rez = 0;

        for current in &what_to_search {
            let now = Instant::now();
            let m = find_three(&ww_in_vec, current);
            rez += now.elapsed().as_nanos();
            what_found.push(m);
        }

        measured_vec.push(rez);

        println!("{:?}", what_found);
    }
    println!("binary :{:?}", measured_binary);
    println!("iter_map :{:?}", measured_iter);
    println!("iter_vec :{:?}", measured_vec);
}

fn find_one(from_what: &BTreeMap<FloatWrapper, f32>, what: &FloatWrapper) -> f32 {
    let v: Vec<f32> = from_what
        .range((Included(what), (Unbounded)))
        .take(1)
        .map(|(_, &v)| v)
        .collect();
    *v.get(0).expect("we are in truble")
}

fn find_two(from_what: &BTreeMap<FloatWrapper, f32>, what: &FloatWrapper) -> f32 {
    from_what
        .iter()
        .skip_while(|(i, _)| *i < what) // Skipping all elements before it
        .take(1) // Reducing the iterator to 1 element
        .map(|(_, &v)| v) // Getting its value, dereferenced
        .next()
        .expect("we are in truble") // Our
}

fn find_three(from_what: &Vec<(FloatWrapper, f32)>, what: &FloatWrapper) -> f32 {
    *from_what
        .iter()
        .skip_while(|(i, _)| i < what) // Skipping all elements before it
        .take(1) // Reducing the iterator to 1 element
        .map(|(_, v)| v) // Getting its value, dereferenced
        .next()
        .expect("we are in truble") // Our
}

The key takeaway for me is that it is worth to use the binary search after ~50 elements. In my case with 30000 elements means 200x speedup (at least based on this microbenchmark).

2 Answers

You said you wanted a std-only solution, but this is a common enough problem, so here's a solution using the crate ordered-float:

Cargo.toml

[dependencies]
ordered-float = "1.0"

main.rs

use ordered_float::OrderedFloat; // 1.0.2
use std::collections::BTreeMap;

fn main() {
    let mut ww = BTreeMap::new();
    ww.insert(OrderedFloat(1.0), "one");
    ww.insert(OrderedFloat(2.0), "two");
    ww.insert(OrderedFloat(3.0), "three");
    ww.insert(OrderedFloat(4.0), "three");
    let rez = ww.range(OrderedFloat(1.5)..).next().map(|(_, &v)| v);

    println!("{:?}", rez);
}

prints

Some("two")

Now, isn't that nice and clean? If you want a less verbose syntax, I suggest wrapping the BTreeMap itself, so you can give it appropriately named methods that make sense for your application.

NaN behavior

Be aware that OrderedFloat may not behave the way you expect in the presence of NaNs:

NaN is sorted as greater than all other values and equal to itself, in contradiction with the IEEE standard.

Now that we've gone over and clarified the requirements a bit, there's a couple of bad news for you:

  1. You're not getting away from the requirement to have a wrapping type. As I'm sure you've discovered, this is because no floating-point type implements Ord
  2. You're also not getting away from a combinator of some sort

First, we're going to clear up your impl, as they both have shortfalls described in the comments. In the future, it may make sense to use the wrapper traits in eq-float, as they already implement all this. The implementations at fault are PartialEq and Ord, and they both break down on a few points. The new implementations:

impl Ord for FloatWrapper {
    fn cmp(&self, other: &Self) -> Ordering {
        self.0.partial_cmp(&other.0).unwrap_or_else(|| {
            if self.0.is_nan() && !other.0.is_nan() {
                Ordering::Less
            } else if !self.0.is_nan() && other.0.is_nan() {
                Ordering::Greater
            } else {
                Ordering::Equal
            }
        })
    }
}
impl PartialEq for FloatWrapper {
    fn eq(&self, other: &Self) -> bool {
        if self.0.is_nan() && other.0.is_nan() {
            true
        } else {
            self.0 == other.0
        }
    }
}

Nothing surprising, we're just abusing the fact that f32 implements PartialOrd for Ord and surfacing all other implementations on FloatWrapper itself.

Now, for the combinator. Your current combinator will force a range of elements to be stored temporarily in memory, to then discard one. We can do better by abusing the fact that iter() is a sorted iterator. So, we can skip while we search, and then take the first:

        let mut first_element = ww.iter()
            .skip_while(|(i, _)| *i < &FloatWrapper::new(1.5)) // Skipping all elements before it
            .take(1) // Reducing the iterator to 1 element
            .map(|(_, &v)| v) // Getting its value, dereferenced
            .next(); // Our result

This yields a 10% speedup in low-element-count situations over your first implementation.

Related