From f5eef18df8c59792e4fd5dda7aebfa1dd6780c87 Mon Sep 17 00:00:00 2001 From: ameligrana Date: Sun, 4 Oct 2026 16:07:10 +0200 Subject: [PATCH 1/2] Fix some minor issues with samplers --- .gitignore | 4 ++ comparison/benchmark_data_results.csv | 13 +++++ comparison/benchmark_results.csv | 37 +++++++++++++ comparison/samplers.jl | 16 +++--- src/SamplingInterface.jl | 6 +-- src/SamplingReduction.jl | 14 ++--- src/SamplingUtils.jl | 1 + src/SortedSamplingMulti.jl | 4 +- src/UnweightedSamplingMulti.jl | 14 ++--- src/UnweightedSamplingSingle.jl | 4 +- src/WeightedSamplingMulti.jl | 40 +++++++------- src/WeightedSamplingSingle.jl | 5 +- test/basic_tests.jl | 16 ++++++ test/merge_tests.jl | 71 ++++++++++++++++++++++++- test/unweighted_sampling_multi_tests.jl | 27 ++++++++++ test/weighted_sampling_multi_tests.jl | 33 +++++++++++- 16 files changed, 251 insertions(+), 54 deletions(-) create mode 100644 comparison/benchmark_data_results.csv create mode 100644 comparison/benchmark_results.csv diff --git a/.gitignore b/.gitignore index 4beb294d..cef4c5ee 100644 --- a/.gitignore +++ b/.gitignore @@ -25,3 +25,7 @@ Manifest.toml benchmark/Manifest.toml benchmark/random_data.arrow + +# Clickstream dataset used by comparison/benchmark_clickstream.jl (download from https://dumps.wikimedia.org/other/clickstream/2024-12/) +comparison/clickstream-*.tsv +comparison/clickstream-*.tsv.gz diff --git a/comparison/benchmark_data_results.csv b/comparison/benchmark_data_results.csv new file mode 100644 index 00000000..70e5b50a --- /dev/null +++ b/comparison/benchmark_data_results.csv @@ -0,0 +1,13 @@ +alg,N,m,repetitions,t_add,t_get +WRSWR-BIN,34240686,3424069,100,28.234017344161856,3.63 +WRSWR-BIN,34240686,342407,100,17.949567220703457,58.800000000000004 +WRSWR-BIN,34240686,34241,100,17.01158814195487,2.56 +WRSWR-BIN,34240686,3425,100,16.926156700248352,4.390000000000001 +WRAExp-J,34240686,3424069,100,63.313972090979725,6.268304051400001e8 +WRAExp-J,34240686,342407,100,7.436388725973536,4.524304526e7 +WRAExp-J,34240686,34241,100,3.5189679575929054,3.76732991e6 +WRAExp-J,34240686,3425,100,3.1451885654393714,349277.97000000003 +WRSWR-SKIP,34240686,3424069,100,8.621851440417988,9.450000000000001 +WRSWR-SKIP,34240686,342407,100,2.5124268161566627,9.22 +WRSWR-SKIP,34240686,34241,100,1.74814293673906,8.180000000000001 +WRSWR-SKIP,34240686,3425,100,1.6209248611432612,8.120000000000001 diff --git a/comparison/benchmark_results.csv b/comparison/benchmark_results.csv new file mode 100644 index 00000000..c91f8c07 --- /dev/null +++ b/comparison/benchmark_results.csv @@ -0,0 +1,37 @@ +alg,N,m,d,repetitions,t_add,t_get +WRSWR-BIN,10000000,1000,d_decr,100,15.351767258000008,3.6700000000000004 +WRSWR-BIN,10000000,10000,d_decr,100,15.476184351000002,3.94 +WRSWR-BIN,10000000,100000,d_decr,100,15.719699089999997,3.66 +WRSWR-BIN,10000000,1000000,d_decr,100,19.063183873999996,4.680000000000001 +WRAExp-J,10000000,1000,d_decr,100,3.0741342810000005,81035.38 +WRAExp-J,10000000,10000,d_decr,100,3.3204783489999996,914356.77 +WRAExp-J,10000000,100000,d_decr,100,5.671898599000001,1.042682492e7 +WRAExp-J,10000000,1000000,d_decr,100,30.680991763000012,1.2113268227000001e8 +WRSWR-SKIP,10000000,1000,d_decr,100,2.5568252619999994,10.120000000000001 +WRSWR-SKIP,10000000,10000,d_decr,100,2.6231644160000007,7.570000000000001 +WRSWR-SKIP,10000000,100000,d_decr,100,3.0551440570000015,9.71 +WRSWR-SKIP,10000000,1000000,d_decr,100,6.853534331000001,9.91 +WRSWR-BIN,10000000,1000,d_const,100,15.514981697000003,4.670000000000001 +WRSWR-BIN,10000000,10000,d_const,100,15.661928006000004,4.65 +WRSWR-BIN,10000000,100000,d_const,100,18.322097050000004,4.7700000000000005 +WRSWR-BIN,10000000,1000000,d_const,100,43.01236574400001,4.700000000000001 +WRAExp-J,10000000,1000,d_const,100,3.1586977939999996,103805.02 +WRAExp-J,10000000,10000,d_const,100,4.0101058410000014,1.06259209e6 +WRAExp-J,10000000,100000,d_const,100,11.213429000999996,1.255766748e7 +WRAExp-J,10000000,1000000,d_const,100,65.902981587,1.5659078689e8 +WRSWR-SKIP,10000000,1000,d_const,100,2.6010059800000014,6.240000000000001 +WRSWR-SKIP,10000000,10000,d_const,100,2.889858204,8.270000000000001 +WRSWR-SKIP,10000000,100000,d_const,100,5.284384767000001,9.72 +WRSWR-SKIP,10000000,1000000,d_const,100,20.140484113000003,9.660000000000002 +WRSWR-BIN,10000000,1000,d_incr,100,15.512825132000005,4.6000000000000005 +WRSWR-BIN,10000000,10000,d_incr,100,15.966000717999995,5.12 +WRSWR-BIN,10000000,100000,d_incr,100,21.065067476000003,6.3 +WRSWR-BIN,10000000,1000000,d_incr,100,66.010967182,3.8900000000000006 +WRAExp-J,10000000,1000,d_incr,100,3.4165506430000008,96650.21 +WRAExp-J,10000000,10000,d_incr,100,4.7180384370000015,1.0320041799999999e6 +WRAExp-J,10000000,100000,d_incr,100,17.500391315000005,1.2637909870000001e7 +WRAExp-J,10000000,1000000,d_incr,100,103.122172538,1.5031790519000003e8 +WRSWR-SKIP,10000000,1000,d_incr,100,2.6298687370000007,7.170000000000001 +WRSWR-SKIP,10000000,10000,d_incr,100,3.216991788000002,6.71 +WRSWR-SKIP,10000000,100000,d_incr,100,7.403542585000004,8.180000000000001 +WRSWR-SKIP,10000000,1000000,d_incr,100,30.919183713000013,9.800000000000002 diff --git a/comparison/samplers.jl b/comparison/samplers.jl index d5c2e576..37fd689f 100644 --- a/comparison/samplers.jl +++ b/comparison/samplers.jl @@ -78,7 +78,7 @@ macro quantile_fast(k) append!(block.args, firstv.args) for i in 2:k nextv = quote - $(esc(:s)) *= ($(esc(:n)) - $i) * $(esc(:p)) + $(esc(:s)) *= ($(esc(:n)) - $(i - 1)) * $(esc(:p)) $(esc(:q)) *= 1. - $(esc(:p)) $(esc(:x)) += $(esc(:s)) / ($(esc(:q)) * $(factorial(i))) $(esc(:x)) > $(esc(:nt)) && return $i @@ -98,7 +98,7 @@ end function get(s::AlgWRSWRSKIPSampler) if s.seen_k < s.n - return sample(s.rng, s.value[1:s.seen_k], Weights(s.weigths[1:s.seen_k]), s.n) + return sample(s.rng, s.value[1:s.seen_k], Weights(diff([0.0; s.weights[1:s.seen_k]])), s.n) else return s.value end @@ -253,17 +253,17 @@ function get(s::AlgWRAExpJSampler{<:BinaryHeap{Pair{T, Tuple{Float64,Float64}}}} m = 2 @inbounds for j in 2:s.n i = rand(s.rng, sampler) - if i <= j-1 + if i < m out[j] = kvs[i][1] else out[j] = kvs[m][1] + if m < s.n + sampler[m] = kvs[m][2][2] + Wnew -= kvs[m][2][2] + sampler[m+1] = Wnew + end m += 1 end - if j < s.n - sampler[j] = kvs[j][2][2] - Wnew -= kvs[j][2][2] - sampler[j+1] = Wnew - end end return out end diff --git a/src/SamplingInterface.jl b/src/SamplingInterface.jl index bebfc3c8..0a783c30 100644 --- a/src/SamplingInterface.jl +++ b/src/SamplingInterface.jl @@ -223,10 +223,10 @@ function itsample(iter, method = AlgRSWRSKIP(); iter_type = infer_eltype(iter)) return itsample(Random.default_rng(), iter, method; iter_type) end function itsample(iter, n::Int, method = AlgL(); iter_type = infer_eltype(iter), ordered = false) - return itsample(Random.default_rng(), iter, n, method; ordered) + return itsample(Random.default_rng(), iter, n, method; iter_type, ordered) end function itsample(iter, wv::Function, method = AlgWRSWRSKIP(); iter_type = infer_eltype(iter)) - return itsample(Random.default_rng(), iter, wv, method) + return itsample(Random.default_rng(), iter, wv, method; iter_type) end function itsample(iter, wv::Function, n::Int, method = AlgAExpJ(); iter_type = infer_eltype(iter), ordered = false) @@ -247,7 +247,7 @@ Base.@constprop :aggressive function itsample(rng::AbstractRNG, iter, n::Int, me s = ReservoirSampler{iter_type,Float64}(rng, n, method, ImmutSampler(), ordered ? Ord() : Unord()) return update_all!(s, iter, ordered) else - m = method isa AlgL || method isa AlgR || method isa AlgD ? AlgD() : AlgORDSWR() + m = method isa Union{AlgL, AlgR, AlgD} ? AlgD() : method isa AlgHiddenShuffle ? method : AlgORDSWR() s = collect(SequentialSampler{iter_type}(rng, iter, n, length(iter), m)) return ordered ? s : fshuffle!(rng, s) end diff --git a/src/SamplingReduction.jl b/src/SamplingReduction.jl index 20b11747..b47ace5c 100644 --- a/src/SamplingReduction.jl +++ b/src/SamplingReduction.jl @@ -3,19 +3,19 @@ const SMWR = Union{MultiAlgRSWRSKIPSampler, MultiAlgWRSWRSKIPSampler} const SMWOWR = Union{MultiAlgAResSampler, MultiAlgAExpJSampler} reduce_samples(t) = error() -function reduce_samples(t::Union{TypeS,TypeUnion}, ss::BinaryHeap...) - nt = length(ss) - n = minimum(length.(ss)) - lkeys = sort(reduce(vcat, [s.valtree for s in ss]), by=(x->x[end]), rev=true)[1:n] - return lkeys +function reduce_samples(t::Union{TypeS,TypeUnion}, n::Integer, ss::BinaryHeap...) + lkeys = sort(reduce(vcat, [s.valtree for s in ss]), by=(x->x[end]), rev=true) + return lkeys[1:min(n, length(lkeys))] end function reduce_samples(ps::AbstractArray, rngs, t::Union{TypeS,TypeUnion}, ss::AbstractArray...) + return reduce_samples(ps, rngs, t, minimum(length.(ss)), ss...) +end +function reduce_samples(ps::AbstractArray, rngs, t::Union{TypeS,TypeUnion}, n::Integer, ss::AbstractArray...) nt = length(ss) T = get_type_rs(t, ss...) v = Vector{Vector{T}}(undef, nt) - n = minimum(length.(ss)) ns = rand(extract_rng(rngs, 1), Multinomial(n, ps)) - Threads.@threads for i in 1:nt + for i in 1:nt s = ss[i] vi = Vector{T}(undef, ns[i]) @inbounds for (q, j) in enumerate(SequentialSampler(extract_rng(rngs, 1), diff --git a/src/SamplingUtils.jl b/src/SamplingUtils.jl index db396308..40f5cb93 100644 --- a/src/SamplingUtils.jl +++ b/src/SamplingUtils.jl @@ -43,6 +43,7 @@ struct SeqSampleIter{R} end @inline function Base.iterate(it::SeqSampleIter) + it.n == 0 && return nothing i = 0 q1 = it.N - it.n + 1 q2 = q1 / it.N diff --git a/src/SortedSamplingMulti.jl b/src/SortedSamplingMulti.jl index a3f6bab6..6f1c200a 100644 --- a/src/SortedSamplingMulti.jl +++ b/src/SortedSamplingMulti.jl @@ -71,7 +71,9 @@ end @inline function Base.iterate(s::MultiAlgORDSampler) indices, iter = s.inds, s.it - curr_idx, state_idx = iterate(indices)::Tuple + it_indices = iterate(indices) + it_indices === nothing && return nothing + curr_idx, state_idx = it_indices el, state_el = iterate(iter)::Tuple for _ in 1:curr_idx-1 el, state_el = iterate(iter, state_el)::Tuple diff --git a/src/UnweightedSamplingMulti.jl b/src/UnweightedSamplingMulti.jl index 3a0d2f58..b693bb50 100644 --- a/src/UnweightedSamplingMulti.jl +++ b/src/UnweightedSamplingMulti.jl @@ -173,7 +173,7 @@ macro quantile_fast(k) append!(block.args, firstv.args) for i in 2:k nextv = quote - $(esc(:s)) *= ($(esc(:n)) - $i) * $(esc(:p)) + $(esc(:s)) *= ($(esc(:n)) - $(i - 1)) * $(esc(:p)) $(esc(:q)) *= 1. - $(esc(:p)) $(esc(:x)) += $(esc(:s)) / ($(esc(:q)) * $(factorial(i))) $(esc(:x)) > $(esc(:nt)) && return $i @@ -216,11 +216,11 @@ function Base.merge(ss::MultiAlgLSampler...) error("To Be Implemented") end function Base.merge(ss::MultiAlgRSWRSKIPSampler...) - newvalue = reduce_samples(get_ps(ss...), [s.rng for s in ss], TypeUnion(), value.(ss)...) - skip_k = sum(getfield(s, :skip_k) for s in ss) - seen_k = sum(getfield(s, :seen_k) for s in ss) n = minimum(s.n for s in ss) - return MultiAlgRSWRSKIPSampler_Mut(n, skip_k, seen_k, ss[1].rng, newvalue, nothing) + newvalue = reduce_samples(get_ps(ss...), [s.rng for s in ss], TypeUnion(), n, value.(ss)...) + seen_k = sum(getfield(s, :seen_k) for s in ss) + s = MultiAlgRSWRSKIPSampler_Mut(n, 0, seen_k, ss[1].rng, newvalue, nothing) + return recompute_skip!(s, n) end function Base.merge!(ss::MultiAlgRSampler...) @@ -231,12 +231,12 @@ function Base.merge!(ss::MultiAlgLSampler...) end function Base.merge!(s1::MultiAlgRSWRSKIPSampler{<:Nothing}, ss::MultiAlgRSWRSKIPSampler...) s1.n > minimum(s.n for s in ss) && error("The size of the mutated reservoir should be the minimum size between all merged reservoir") - newvalue = reduce_samples(get_ps(s1, ss...), [s1.rng, [s.rng for s in ss]...], TypeS(), value(s1), value.(ss)...) + newvalue = reduce_samples(get_ps(s1, ss...), [s1.rng, [s.rng for s in ss]...], TypeS(), s1.n, value(s1), value.(ss)...) for i in 1:length(newvalue) @inbounds s1.value[i] = newvalue[i] end - s1.skip_k += sum(getfield(s, :skip_k) for s in ss) s1.seen_k += sum(getfield(s, :seen_k) for s in ss) + recompute_skip!(s1, s1.n) return s1 end diff --git a/src/UnweightedSamplingSingle.jl b/src/UnweightedSamplingSingle.jl index 318fb712..aba795d1 100644 --- a/src/UnweightedSamplingSingle.jl +++ b/src/UnweightedSamplingSingle.jl @@ -46,7 +46,7 @@ function Base.merge(ss::SingleAlgRSWRSKIPSampler...) ps = cumsum(ns ./ n_tot) r = rand(ss[1].rng) value = ss[findfirst(p -> r < p, ps)].rvalue - return typeof(ss[1])(n_tot, sum(s.skip_k for s in ss), ss[1].rng, value) + return typeof(ss[1])(n_tot, ceil(Int, n_tot/rand(ss[1].rng)), ss[1].rng, value) end function Base.merge!(s1::SingleAlgRSWRSKIPSampler_Mut, ss::SingleAlgRSWRSKIPSampler_Mut...) @@ -59,6 +59,6 @@ function Base.merge!(s1::SingleAlgRSWRSKIPSampler_Mut, ss::SingleAlgRSWRSKIPSamp s1.rvalue = RefVal_Immut(ss[i-1].rvalue.value) end s1.seen_k += sum(s.seen_k for s in ss) - s1.skip_k += sum(s.skip_k for s in ss) + s1.skip_k = ceil(Int, s1.seen_k/rand(s1.rng)) return s1 end diff --git a/src/WeightedSamplingMulti.jl b/src/WeightedSamplingMulti.jl index 7f1998a6..54bd5cba 100644 --- a/src/WeightedSamplingMulti.jl +++ b/src/WeightedSamplingMulti.jl @@ -137,7 +137,7 @@ end while s.weights[j] < curx * s.state j += 1 end - newvalues[i] = s.value[j] + newvalues[n-i+1] = s.value[j] end s.value .= newvalues s = @inline recompute_skip!(s, n) @@ -189,38 +189,37 @@ end extract_T(::DataStructures.BinaryHeap{T}) where T = T function Base.merge(ss::MultiAlgAResSampler...) - newvalue = reduce_samples(TypeUnion(), [s.value for s in ss]...) + n = minimum(s.n for s in ss) + newvalue = reduce_samples(TypeUnion(), n, [s.value for s in ss]...) newheap = BinaryHeap(Base.By(last, DataStructures.FasterForward()), newvalue) seen_k = sum(getfield(s, :seen_k) for s in ss) - n = minimum(s.n for s in ss) s = MultiAlgAResSampler_Mut(seen_k, n, ss[1].rng, newheap) return s end function Base.merge(ss::MultiAlgAExpJSampler...) - newvalue = reduce_samples(TypeUnion(), [s.value for s in ss]...) + n = minimum(s.n for s in ss) + newvalue = reduce_samples(TypeUnion(), n, [s.value for s in ss]...) newheap = BinaryHeap(Base.By(last, DataStructures.FasterForward()), newvalue) seen_k = sum(getfield(s, :seen_k) for s in ss) - state = sum(getfield(s, :state) for s in ss) - min_priority = minimum(getfield(s, :min_priority) for s in ss) - n = minimum(s.n for s in ss) - s = MultiAlgAExpJSampler_Mut(state, min_priority, seen_k, n, ss[1].rng, newheap) + z = zero(ss[1].state) + s = MultiAlgAExpJSampler_Mut(z, z, seen_k, n, ss[1].rng, newheap) + seen_k >= n && recompute_skip!(s) return s end function Base.merge(ss::MultiAlgWRSWRSKIPSampler...) - newvalue = reduce_samples(get_ps(ss...), [s.rng for s in ss], TypeUnion(), value.(ss)...) - skip_w = sum(getfield(s, :skip_w) for s in ss) + n = minimum(s.n for s in ss) + newvalue = reduce_samples(get_ps(ss...), [s.rng for s in ss], TypeUnion(), n, value.(ss)...) state = sum(getfield(s, :state) for s in ss) seen_k = sum(getfield(s, :seen_k) for s in ss) - n = minimum(s.n for s in ss) - s = MultiAlgWRSWRSKIPSampler_Mut(n, state, skip_w, seen_k, ss[1].rng, Memory{Float64}(undef,0), newvalue, nothing) - return s + s = MultiAlgWRSWRSKIPSampler_Mut(n, state, zero(state), seen_k, ss[1].rng, Memory{Float64}(undef,0), newvalue, nothing) + return recompute_skip!(s, n) end function Base.merge!(s1::MultiAlgAResSampler, ss::MultiAlgAResSampler...) - length(typeof(s1.value.valtree).parameters) == 3 && error("Merging ordered reservoirs is not possible") + eltype(s1.value.valtree) <: Tuple && error("Merging ordered reservoirs is not possible") s1.n > minimum(s.n for s in ss) && error("The size of the mutated reservoir should be the minimum size between all merged reservoir") + newvalue = reduce_samples(TypeS(), s1.n, s1.value, [s.value for s in ss]...) empty!(s1.value.valtree) - newvalue = reduce_samples(TypeS(), s1.value, [s.value for s in ss]...) for e in newvalue push!(s1.value, e[1] => e[2]) end @@ -228,27 +227,26 @@ function Base.merge!(s1::MultiAlgAResSampler, ss::MultiAlgAResSampler...) return s1 end function Base.merge!(s1::MultiAlgAExpJSampler, ss::MultiAlgAExpJSampler...) - length(typeof(s1.value.valtree).parameters) == 3 && error("Merging ordered reservoirs is not possible") + eltype(s1.value.valtree) <: Tuple && error("Merging ordered reservoirs is not possible") s1.n > minimum(s.n for s in ss) && error("The size of the mutated reservoir should be the minimum size between all merged reservoir") + newvalue = reduce_samples(TypeS(), s1.n, s1.value, [s.value for s in ss]...) empty!(s1.value.valtree) - newvalue = reduce_samples(TypeS(), s1.value, [s.value for s in ss]...) for e in newvalue push!(s1.value, e[1] => e[2]) end s1.seen_k += sum(getfield(s, :seen_k) for s in ss) - s1.state += sum(getfield(s, :state) for s in ss) - s1.min_priority = min(s1.min_priority, minimum(getfield(s, :min_priority) for s in ss)) + s1.seen_k >= s1.n && recompute_skip!(s1) return s1 end function Base.merge!(s1::MultiAlgWRSWRSKIPSampler{<:Nothing}, ss::MultiAlgWRSWRSKIPSampler...) s1.n > minimum(s.n for s in ss) && error("The size of the mutated reservoir should be the minimum size between all merged reservoir") - newvalue = reduce_samples(get_ps(s1, ss...), [s1.rng, [s.rng for s in ss]...], TypeS(), value(s1), value.(ss)...) + newvalue = reduce_samples(get_ps(s1, ss...), [s1.rng, [s.rng for s in ss]...], TypeS(), s1.n, value(s1), value.(ss)...) for i in 1:length(newvalue) @inbounds s1.value[i] = newvalue[i] end - s1.skip_w += sum(getfield(s, :skip_w) for s in ss) s1.state += sum(getfield(s, :state) for s in ss) s1.seen_k += sum(getfield(s, :seen_k) for s in ss) + recompute_skip!(s1, s1.n) return s1 end diff --git a/src/WeightedSamplingSingle.jl b/src/WeightedSamplingSingle.jl index 5ac3ef96..13996db7 100644 --- a/src/WeightedSamplingSingle.jl +++ b/src/WeightedSamplingSingle.jl @@ -51,8 +51,7 @@ function Base.merge(ss::SingleAlgWRSWRSKIPSampler...) ps = cumsum(ns ./ n_tot) r = rand(ss[1].rng) value = ss[findfirst(p -> r < p, ps)].rvalue - return typeof(ss[1])(sum(s.seen_k for s in ss), sum(s.total_w for s in ss), sum(s.skip_w for s in ss), - ss[1].rng, value) + return typeof(ss[1])(sum(s.seen_k for s in ss), n_tot, n_tot/rand(ss[1].rng), ss[1].rng, value) end function Base.merge!(s1::SingleAlgWRSWRSKIPSampler_Mut, ss::SingleAlgWRSWRSKIPSampler_Mut...) @@ -65,8 +64,8 @@ function Base.merge!(s1::SingleAlgWRSWRSKIPSampler_Mut, ss::SingleAlgWRSWRSKIPSa s1.rvalue = RefVal_Immut(ss[i-1].rvalue.value) end s1.seen_k += sum(s.seen_k for s in ss) - s1.skip_w += sum(s.skip_w for s in ss) s1.total_w += sum(s.total_w for s in ss) + s1.skip_w = s1.total_w/rand(s1.rng) return s1 end diff --git a/test/basic_tests.jl b/test/basic_tests.jl index d6005c14..b3fd9855 100644 --- a/test/basic_tests.jl +++ b/test/basic_tests.jl @@ -3,4 +3,20 @@ @test 1 <= sum(SequentialSampler(1, 10)) <= 10 @test 1 <= sum(SequentialSampler(1, 10, AlgHiddenShuffle())) <= 10 @test length(combine([[1,2,3], [4,5]], [1.0, 2.0])) == 2 +end +@testset "itsample interface" begin + rng = StableRNG(55) + @test all(1:1000) do _ + s = itsample(rng, 1:10, 5, AlgHiddenShuffle()) + length(s) == 5 && allunique(s) + end + @test issorted(itsample(rng, 1:10, 5, AlgHiddenShuffle(); ordered = true)) + for alg in (AlgD(), AlgHiddenShuffle(), AlgORDSWR()) + @test isempty(collect(SequentialSampler(rng, 0, 10, alg))) + @test collect(SequentialSampler{Int}(rng, 1:10, 0, 10, alg)) == Int[] + end + @test itsample(rng, 1:10, 0) == Int[] + iter = Iterators.filter(isodd, 1:10) + @test itsample(iter, 2; iter_type = Float64) isa Vector{Float64} + @test itsample(iter, x -> 1.0; iter_type = Float64) isa Float64 end \ No newline at end of file diff --git a/test/merge_tests.jl b/test/merge_tests.jl index 130ff827..7dec83df 100644 --- a/test/merge_tests.jl +++ b/test/merge_tests.jl @@ -1,6 +1,6 @@ @testset "merge/merge! tests" begin - rng = StableRNG(44) + rng = StableRNG(47) iters = (1:2, 3:10) reps = 10^5 size = 2 @@ -93,3 +93,72 @@ @test pvalue(chisq_test) > 0.05 end end + +is_weighted(m) = m isa Union{AlgWRSWRSKIP, AlgARes, AlgAExpJ} +upd!(s, m, x) = is_weighted(m) ? fit!(s, x, 1.0) : fit!(s, x) + +@testset "fit! after merge/merge!" begin + rng = StableRNG(53) + reps = 10^5 + iters, rest, N = (1:5, 6:10), 11:15, 15 + function merged_sampler(f, m, args...) + s1 = ReservoirSampler{Int}(rng, args..., m) + s2 = ReservoirSampler{Int}(rng, args..., m) + for x in iters[1] upd!(s1, m, x) end + for x in iters[2] upd!(s2, m, x) end + s = f(s1, s2) + for x in rest upd!(s, m, x) end + return s + end + for f in (merge, merge!) + for m in (AlgRSWRSKIP(), AlgWRSWRSKIP()) + counts = zeros(Int, N) + for _ in 1:reps + counts[value(merged_sampler(f, m))] += 1 + end + @test pvalue(ChisqTest(counts, fill(1/N, N))) > 0.001 + counts = zeros(Int, N) + for _ in 1:reps + for x in value(merged_sampler(f, m, 2)) counts[x] += 1 end + end + @test pvalue(ChisqTest(counts, fill(1/N, N))) > 0.001 + end + for m in (AlgARes(), AlgAExpJ()) + counts = Dict{Vector{Int}, Int}() + for _ in 1:reps + k = sort(value(merged_sampler(f, m, 2))) + counts[k] = get(counts, k, 0) + 1 + end + pairs_all = [[i, j] for i in 1:N for j in i+1:N] + count_est = [get(counts, k, 0) for k in pairs_all] + @test pvalue(ChisqTest(count_est, fill(1/length(pairs_all), length(pairs_all)))) > 0.001 + end + end +end + +@testset "merge/merge! keep the reservoir size" begin + rng = StableRNG(54) + for m in (AlgARes(), AlgAExpJ()), f in (merge, merge!) + s1, s2 = ReservoirSampler{Int}(rng, 3, m), ReservoirSampler{Int}(rng, 3, m) + fit!(s1, 1, 1.0) + for x in 2:10 fit!(s2, x, 1.0) end + v = value(f(s1, s2)) + @test length(v) == 3 && allunique(v) && all(in(1:10), v) + s1, s2 = ReservoirSampler{Int}(rng, 3, m), ReservoirSampler{Int}(rng, 3, m) + fit!(s1, 1, 1.0) + fit!(s2, 2, 1.0) + @test sort(value(f(s1, s2))) == [1, 2] + end + for m in (AlgRSWRSKIP(), AlgWRSWRSKIP()), f in (merge, merge!) + s1, s2 = ReservoirSampler{Int}(rng, 3, m), ReservoirSampler{Int}(rng, 3, m) + for x in 1:10 upd!(s2, m, x) end + v = value(f(s1, s2)) + @test length(v) == 3 && all(in(1:10), v) + end + for m in (AlgARes(), AlgAExpJ()) + s1, s2 = ReservoirSampler{Int}(rng, 2, m; ordered = true), ReservoirSampler{Int}(rng, 2, m; ordered = true) + for x in 1:5 fit!(s1, x, 1.0) end + for x in 6:10 fit!(s2, x, 1.0) end + @test_throws "Merging ordered reservoirs is not possible" merge!(s1, s2) + end +end diff --git a/test/unweighted_sampling_multi_tests.jl b/test/unweighted_sampling_multi_tests.jl index 0ea47f1c..d5870253 100644 --- a/test/unweighted_sampling_multi_tests.jl +++ b/test/unweighted_sampling_multi_tests.jl @@ -84,3 +84,30 @@ end end end + +@testset "Unweighted skip sampling with n >= 4" begin + rng = StableRNG(49) + # slots replaced at a skip event follow a Binomial(n, p) conditioned on being positive + for (n, p) in ((4, 0.25), (8, 0.4), (20, 0.1)) + d = Binomial(n, p) + reps = 10^6 + kmax = findlast(k -> reps * pdf(d, k) / ccdf(d, 0) >= 100, 1:n) + counts = zeros(Int, kmax) + for _ in 1:reps + counts[min(StreamSampling.choose(rng, n, p), kmax)] += 1 + end + ps = [[pdf(d, k) for k in 1:kmax-1]; ccdf(d, kmax-1)] ./ ccdf(d, 0) + @test pvalue(ChisqTest(counts, ps)) > 0.001 + end + # each slot of a with-replacement reservoir is an independent uniform draw + N, n, reps = 20, 10, 10^5 + for ordered in (false, true) + counts = zeros(Int, N) + for _ in 1:reps + rs = ReservoirSampler{Int}(rng, n, AlgRSWRSKIP(); ordered) + for x in 1:N fit!(rs, x) end + for x in value(rs) counts[x] += 1 end + end + @test pvalue(ChisqTest(counts, fill(1/N, N))) > 0.001 + end +end diff --git a/test/weighted_sampling_multi_tests.jl b/test/weighted_sampling_multi_tests.jl index a95f1ff5..91252723 100644 --- a/test/weighted_sampling_multi_tests.jl +++ b/test/weighted_sampling_multi_tests.jl @@ -78,7 +78,7 @@ end weight2(el) = el <= 5 ? 1.0 : 2.0 weight3(el) = el <= 5 ? 1.0 : 2.0 wfuncs = (weight2, weight3) - rngs = (StableRNG(41), StableRNG(42)) + rngs = (StableRNG(57), StableRNG(58)) iters = (a:b, Iterators.filter(x -> x != b+1, a:b+1)) sizes = (1, 2) for it in iters @@ -117,3 +117,34 @@ end end end end + +@testset "Weighted skip sampling with n >= 4" begin + rng = StableRNG(51) + # each slot of a with-replacement reservoir is an independent weighted draw + N, n, reps = 20, 10, 10^5 + w = collect(1.0:N) + for ordered in (false, true) + counts = zeros(Int, N) + for _ in 1:reps + rs = ReservoirSampler{Int}(rng, n, AlgWRSWRSKIP(); ordered) + for x in 1:N fit!(rs, x, w[x]) end + for x in value(rs) counts[x] += 1 end + end + @test pvalue(ChisqTest(counts, w ./ sum(w))) > 0.001 + end +end + +@testset "Weighted ordered values follow the stream order" begin + rng = StableRNG(52) + stream = [50, 10, 40, 20, 30, 60, 0] + pos = Dict(x => i for (i, x) in enumerate(stream)) + for method in (AlgARes(), AlgAExpJ(), AlgWRSWRSKIP()), n in (4, 7, 10) + in_order = true + for _ in 1:1000 + rs = ReservoirSampler{Int}(rng, n, method; ordered = true) + for x in stream fit!(rs, x, 1.0) end + in_order &= issorted(ordvalue(rs), by = x -> pos[x]) + end + @test in_order + end +end From a5bb796b8176f0915a87c5f43308fb21c3f35c09 Mon Sep 17 00:00:00 2001 From: ameligrana Date: Mon, 5 Oct 2026 18:22:11 +0200 Subject: [PATCH 2/2] Merge reservoirs with replacement exactly when they are not yet full When the merged reservoirs saw fewer than n elements in total, they are all still filling up. merge/merge! now concatenate their raw elements (with shifted cumulative weights for AlgWRSWRSKIP) instead of drawing n samples, which biased value() and made AlgWRSWRSKIP throw. --- src/SamplingReduction.jl | 23 +++++++++++++++++++++++ src/UnweightedSamplingMulti.jl | 22 ++++++++++++++++------ src/WeightedSamplingMulti.jl | 25 ++++++++++++++++++------- test/merge_tests.jl | 31 +++++++++++++++++++++++++++++++ 4 files changed, 88 insertions(+), 13 deletions(-) diff --git a/src/SamplingReduction.jl b/src/SamplingReduction.jl index b47ace5c..89a6850e 100644 --- a/src/SamplingReduction.jl +++ b/src/SamplingReduction.jl @@ -27,6 +27,29 @@ function reduce_samples(ps::AbstractArray, rngs, t::Union{TypeS,TypeUnion}, n::I return reduce(vcat, v) end +# reservoirs which together saw fewer than n elements are all still filling up, so +# they hold their elements as seen and merging them amounts to concatenating these +function append_unfilled!(value, k, ss::MultiAlgRSWRSKIPSampler...) + for s in ss + @inbounds for i in 1:s.seen_k + value[k+i] = s.value[i] + end + k += s.seen_k + end + return value +end +function append_unfilled!(value, weights, k, w, ss::MultiAlgWRSWRSKIPSampler...) + for s in ss + @inbounds for i in 1:s.seen_k + value[k+i] = s.value[i] + weights[k+i] = w + s.weights[i] + end + k += s.seen_k + w += s.state + end + return value +end + extract_rng(v::AbstractArray, i) = v[i] extract_rng(v::AbstractRNG, i) = v diff --git a/src/UnweightedSamplingMulti.jl b/src/UnweightedSamplingMulti.jl index b693bb50..0ec02865 100644 --- a/src/UnweightedSamplingMulti.jl +++ b/src/UnweightedSamplingMulti.jl @@ -217,8 +217,13 @@ function Base.merge(ss::MultiAlgLSampler...) end function Base.merge(ss::MultiAlgRSWRSKIPSampler...) n = minimum(s.n for s in ss) - newvalue = reduce_samples(get_ps(ss...), [s.rng for s in ss], TypeUnion(), n, value.(ss)...) seen_k = sum(getfield(s, :seen_k) for s in ss) + if seen_k < n + newvalue = Vector{get_type_rs(TypeUnion(), (s.value for s in ss)...)}(undef, n) + append_unfilled!(newvalue, 0, ss...) + return MultiAlgRSWRSKIPSampler_Mut(n, 0, seen_k, ss[1].rng, newvalue, nothing) + end + newvalue = reduce_samples(get_ps(ss...), [s.rng for s in ss], TypeUnion(), n, value.(ss)...) s = MultiAlgRSWRSKIPSampler_Mut(n, 0, seen_k, ss[1].rng, newvalue, nothing) return recompute_skip!(s, n) end @@ -231,12 +236,17 @@ function Base.merge!(ss::MultiAlgLSampler...) end function Base.merge!(s1::MultiAlgRSWRSKIPSampler{<:Nothing}, ss::MultiAlgRSWRSKIPSampler...) s1.n > minimum(s.n for s in ss) && error("The size of the mutated reservoir should be the minimum size between all merged reservoir") - newvalue = reduce_samples(get_ps(s1, ss...), [s1.rng, [s.rng for s in ss]...], TypeS(), s1.n, value(s1), value.(ss)...) - for i in 1:length(newvalue) - @inbounds s1.value[i] = newvalue[i] + seen_k = s1.seen_k + sum(getfield(s, :seen_k) for s in ss) + if seen_k < s1.n + append_unfilled!(s1.value, s1.seen_k, ss...) + else + newvalue = reduce_samples(get_ps(s1, ss...), [s1.rng, [s.rng for s in ss]...], TypeS(), s1.n, value(s1), value.(ss)...) + for i in 1:length(newvalue) + @inbounds s1.value[i] = newvalue[i] + end end - s1.seen_k += sum(getfield(s, :seen_k) for s in ss) - recompute_skip!(s1, s1.n) + s1.seen_k = seen_k + seen_k >= s1.n && recompute_skip!(s1, s1.n) return s1 end diff --git a/src/WeightedSamplingMulti.jl b/src/WeightedSamplingMulti.jl index 54bd5cba..c8df3f16 100644 --- a/src/WeightedSamplingMulti.jl +++ b/src/WeightedSamplingMulti.jl @@ -208,10 +208,16 @@ function Base.merge(ss::MultiAlgAExpJSampler...) end function Base.merge(ss::MultiAlgWRSWRSKIPSampler...) n = minimum(s.n for s in ss) - newvalue = reduce_samples(get_ps(ss...), [s.rng for s in ss], TypeUnion(), n, value.(ss)...) state = sum(getfield(s, :state) for s in ss) seen_k = sum(getfield(s, :seen_k) for s in ss) - s = MultiAlgWRSWRSKIPSampler_Mut(n, state, zero(state), seen_k, ss[1].rng, Memory{Float64}(undef,0), newvalue, nothing) + weights = Memory{typeof(state)}(undef, n) + if seen_k < n + newvalue = Vector{get_type_rs(TypeUnion(), (s.value for s in ss)...)}(undef, n) + append_unfilled!(newvalue, weights, 0, zero(state), ss...) + return MultiAlgWRSWRSKIPSampler_Mut(n, state, zero(state), seen_k, ss[1].rng, weights, newvalue, nothing) + end + newvalue = reduce_samples(get_ps(ss...), [s.rng for s in ss], TypeUnion(), n, value.(ss)...) + s = MultiAlgWRSWRSKIPSampler_Mut(n, state, zero(state), seen_k, ss[1].rng, weights, newvalue, nothing) return recompute_skip!(s, n) end @@ -240,13 +246,18 @@ function Base.merge!(s1::MultiAlgAExpJSampler, ss::MultiAlgAExpJSampler...) end function Base.merge!(s1::MultiAlgWRSWRSKIPSampler{<:Nothing}, ss::MultiAlgWRSWRSKIPSampler...) s1.n > minimum(s.n for s in ss) && error("The size of the mutated reservoir should be the minimum size between all merged reservoir") - newvalue = reduce_samples(get_ps(s1, ss...), [s1.rng, [s.rng for s in ss]...], TypeS(), s1.n, value(s1), value.(ss)...) - for i in 1:length(newvalue) - @inbounds s1.value[i] = newvalue[i] + seen_k = s1.seen_k + sum(getfield(s, :seen_k) for s in ss) + if seen_k < s1.n + append_unfilled!(s1.value, s1.weights, s1.seen_k, s1.state, ss...) + else + newvalue = reduce_samples(get_ps(s1, ss...), [s1.rng, [s.rng for s in ss]...], TypeS(), s1.n, value(s1), value.(ss)...) + for i in 1:length(newvalue) + @inbounds s1.value[i] = newvalue[i] + end end s1.state += sum(getfield(s, :state) for s in ss) - s1.seen_k += sum(getfield(s, :seen_k) for s in ss) - recompute_skip!(s1, s1.n) + s1.seen_k = seen_k + seen_k >= s1.n && recompute_skip!(s1, s1.n) return s1 end diff --git a/test/merge_tests.jl b/test/merge_tests.jl index 7dec83df..2eabe76e 100644 --- a/test/merge_tests.jl +++ b/test/merge_tests.jl @@ -162,3 +162,34 @@ end @test_throws "Merging ordered reservoirs is not possible" merge!(s1, s2) end end + +@testset "merge/merge! below the reservoir size" begin + rng = StableRNG(56) + reps = 10^5 + # weights equal to the elements, so that misplaced weights are detected + fitx!(s, m, x) = m isa AlgWRSWRSKIP ? fit!(s, x, Float64(x)) : fit!(s, x) + for m in (AlgRSWRSKIP(), AlgWRSWRSKIP()), f in (merge, merge!), N in (4, 6, 10) + counts = zeros(Int, N) + for _ in 1:reps + s1, s2 = ReservoirSampler{Int}(rng, 6, m), ReservoirSampler{Int}(rng, 6, m) + for x in 1:2 fitx!(s1, m, x) end + for x in 3:4 fitx!(s2, m, x) end + s = f(s1, s2) + for x in 5:N fitx!(s, m, x) end + for x in value(s) counts[x] += 1 end + end + ps = m isa AlgWRSWRSKIP ? collect(1:N) ./ sum(1:N) : fill(1/N, N) + @test pvalue(ChisqTest(counts, ps)) > 0.001 + end + for m in (AlgRSWRSKIP(), AlgWRSWRSKIP()), f in (merge, merge!) + s = f(ReservoirSampler{Int}(rng, 3, m), ReservoirSampler{Int}(rng, 3, m)) + @test isempty(value(s)) + fitx!(s, m, 1) + @test value(s) == [1, 1, 1] + s1, s2 = ReservoirSampler{Int}(rng, 3, m), ReservoirSampler{Int}(rng, 3, m) + for x in 1:5 fitx!(s1, m, x) end + s = empty!(f(s1, s2)) + fitx!(s, m, 1) + @test value(s) == [1, 1, 1] + end +end