function check_k(tree, k) if k > length(get_tree(tree).data) || k < 0 throw(ArgumentError("k > number of points in tree or < 0")) end end """ allnn(tree::NNTree [, skip=always_false]) -> indices, distances Compute the nearest neighbor for every point stored in `tree`, excluding each point itself. Returns two vectors of length `npoints` containing the neighbor index and distance for each point. """ function allnn(tree::NNTree{V}, skip::F=Returns(false)) where {V, F<:Function} check_valid(tree) inner_tree = get_tree(tree) n_points = length(inner_tree.data) n_points == 0 && return Vector{Int}(), Vector{get_T(eltype(V))}() n_points == 1 && throw(ArgumentError("allnn requires at least 2 points")) idxs = Vector{Int}(undef, n_points) dists = Vector{get_T(eltype(V))}(undef, n_points) for i in 1:n_points orig_idx = inner_tree.reordered ? inner_tree.indices[i] : i best_idx, best_dist = _knn(tree, inner_tree.data[i], -1, dist_typemax(inner_tree), nothing, skip, orig_idx) best_idx == -1 && throw(ArgumentError("no neighbor found for point $orig_idx: all candidate points were skipped")) idxs[orig_idx] = inner_tree.reordered ? inner_tree.indices[best_idx] : best_idx dists[orig_idx] = best_dist end return idxs, dists end """ allknn(tree::NNTree, k [, sortres=false, skip=always_false]) -> indices, distances Compute the `k` nearest neighbors for every point stored in `tree`, excluding each point itself. Returns two vectors of length `npoints`, each containing a length-`k` vector of neighbor indices and distances, respectively. Set `sortres=true` to order neighbors by distance. """ function allknn(tree::NNTree{V}, k::Int, sortres=false, skip::F=Returns(false)) where {V, F<:Function} check_valid(tree) inner_tree = get_tree(tree) n_points = length(inner_tree.data) n_points == 0 && return Vector{Vector{Int}}(), Vector{Vector{get_T(eltype(V))}}() k < 0 && throw(ArgumentError("k < 0")) k <= n_points - 1 || throw(ArgumentError("k must be <= number of points - 1 for allknn")) dists = [Vector{get_T(eltype(V))}(undef, k) for _ in 1:n_points] idxs = [Vector{Int}(undef, k) for _ in 1:n_points] for i in 1:n_points orig_idx = inner_tree.reordered ? inner_tree.indices[i] : i knn_point!(tree, inner_tree.data[i], sortres, dists[orig_idx], idxs[orig_idx], skip, orig_idx) end return idxs, dists end """ knn(tree::NNTree, points, k [, skip=always_false]) -> indices, distances Performs a lookup of the `k` nearest neighbors to the `points` from the data in the `tree`. # Arguments - `tree`: The tree instance - `points`: Query point(s) - can be a vector (single point), matrix (multiple points), or vector of vectors - `k`: Number of nearest neighbors to find - `skip`: Optional predicate function to skip points based on their index (default: `always_false`) # Returns - `indices`: Indices of the k nearest neighbors - `distances`: Distances to the k nearest neighbors See also: `knn!`, `nn`. """ function knn(tree::NNTree{V}, points::AbstractVector{T}, k::Int, sortres=false, skip::F=Returns(false)) where {V, T <: AbstractVector, F<:Function} check_input(tree, points) check_for_nan_in_points(points) check_k(tree, k) n_points = length(points) dists = [Vector{get_T(eltype(V))}(undef, k) for _ in 1:n_points] idxs = [Vector{Int}(undef, k) for _ in 1:n_points] for i in 1:n_points knn_point!(tree, points[i], sortres, dists[i], idxs[i], skip) end return idxs, dists end knn_point!(tree::NNTree{V}, point::AbstractVector{T}, sortres, dist, idx, skip::F, self_idx::Int=0) where {V, T <: Number, F} = _knn_point!(tree, point, sortres, dist, idx, skip, self_idx) function _knn_point!(tree::NNTree{V}, point::AbstractVector{T}, sortres, dist_final, idx, skip::F, self_idx::Int) where {V, T <: Number, F} isempty(idx) && return # k == 0 fill!(idx, -1) inner_tree = get_tree(tree) T_internal = dist_type_internal(inner_tree) T_final = eltype(dist_final) if T_internal === T_final dist_internal = dist_final else dist_internal = Vector{T_internal}(undef, length(dist_final)) end fill!(dist_internal, dist_typemax(inner_tree)) _, ret_dists = _knn(tree, point, idx, dist_internal, dist_final, skip, self_idx) # Trees that finalize distances themselves (KDTree) return `dist_final`; # for the others convert the internal distances into the output vector. if ret_dists !== dist_final copyto!(dist_final, ret_dists) end if skip !== Returns(false) # Compact away unfilled entries (k larger than the number of non-skipped points) j = 0 @inbounds for t in eachindex(idx) if idx[t] != -1 j += 1 idx[j] = idx[t] dist_final[j] = dist_final[t] end end resize!(idx, j) resize!(dist_final, j) end sortres && heap_sort_inplace!(dist_final, idx) if inner_tree.reordered for j in eachindex(idx) @inbounds idx[j] = inner_tree.indices[idx[j]] end end return end """ knn!(idxs, dists, tree, point, k [, skip=always_false]) Same functionality as `knn` but stores the results in the input vectors `idxs` and `dists`. Useful to avoid allocations or specify the element type of the output vectors. # Arguments - `idxs`: Pre-allocated vector to store indices (must be of length `k`) - `dists`: Pre-allocated vector to store distances (must be of length `k`); the element type must be able to represent the computed distances (e.g. `Float32` works for a `Float64` tree) - `tree`: The tree instance - `point`: Query point - `k`: Number of nearest neighbors to find - `skip`: Optional predicate function to skip points based on their index (default: `always_false`) See also: `knn`, `nn`. """ function knn!(idxs::AbstractVector{<:Integer}, dists::AbstractVector, tree::NNTree{V}, point::AbstractVector{T}, k::Int, sortres=false, skip::F=Returns(false)) where {V, T <: Number, F<:Function} check_input(tree, point) check_for_nan_in_points(point) check_k(tree, k) length(idxs) == k || throw(ArgumentError("idxs must be of length k")) length(dists) == k || throw(ArgumentError("dists must be of length k")) knn_point!(tree, point, sortres, dists, idxs, skip) return idxs, dists end function knn(tree::NNTree{V}, point::AbstractVector{T}, k::Int, sortres=false, skip::F=Returns(false)) where {V, T <: Number, F<:Function} idx = Vector{Int}(undef, k) dist = Vector{get_T(eltype(V))}(undef, k) return knn!(idx, dist, tree, point, k, sortres, skip) end function knn(tree::NNTree{V}, points::AbstractMatrix{T}, k::Int, sortres=false, skip::F=Returns(false)) where {V, T <: Number, F<:Function} dim = size(points, 1) knn_matrix(tree, points, k, Val(dim), sortres, skip) end function knn_matrix(tree::NNTree{V}, points::AbstractMatrix{T}, k::Int, ::Val{dim}, sortres=false, skip::F=Returns(false)) where {V, T <: Number, dim, F<:Function} # TODO: DRY with knn for AbstractVector check_input(tree, points) check_for_nan_in_points(points) check_k(tree, k) n_points = size(points, 2) dists = [Vector{get_T(eltype(V))}(undef, k) for _ in 1:n_points] idxs = [Vector{Int}(undef, k) for _ in 1:n_points] for i in 1:n_points point = SVector{dim,T}(ntuple(j -> points[j, i], Val(dim))) knn_point!(tree, point, sortres, dists[i], idxs[i], skip) end return idxs, dists end """ nn(tree::NNTree, point [, skip]) -> index, distance nn(tree::NNTree, points [, skip]) -> indices, distances Performs a lookup of the single nearest neighbor to the `point(s)` from the data. # Arguments - `tree`: The tree instance - `point(s)`: Query point(s) - can be a vector (single point), matrix (multiple points), or vector of vectors - `skip`: Optional predicate function to skip points based on their index (default: `always_false`) # Returns - For single point: `index` and `distance` of the nearest neighbor - For multiple points: vectors of `indices` and `distances` of the nearest neighbors See also: `knn`. """ function nn(tree::NNTree{V}, point::AbstractVector{T}, skip::F=Returns(false)) where {V, T <: Number, F <: Function} check_input(tree, point) check_for_nan_in_points(point) check_k(tree, 1) best_idx, best_dist = _knn(tree, point, -1, dist_typemax(get_tree(tree)), nothing, skip, 0) best_idx == -1 && throw(ArgumentError("no neighbor found: all points in the tree were skipped")) inner_tree = get_tree(tree) final_idx = inner_tree.reordered ? inner_tree.indices[best_idx] : best_idx return final_idx, best_dist end nn(tree::NNTree{V}, points::AbstractVector{T}, skip::F=Returns(false)) where {V, T <: AbstractVector, F <: Function} = _nn(tree, points, skip) |> _onlyeach nn(tree::NNTree{V}, points::AbstractMatrix{T}, skip::F=Returns(false)) where {V, T <: Number, F <: Function} = _nn(tree, points, skip) |> _onlyeach _nn(tree, points, skip) = knn(tree, points, 1, false, skip) _onlyeach(v::Tuple) = only.(first(v)), only.(last(v))