sampling: partial top-p select in dist_build — O(V) heapify + k pops, not O(V log V) qsort (#335)
dist_build() sorted the entire 151936-entry vocab by probability (qsort) on every sampled token whenever 0 < g_nuc < 1 — the serve default — and again per draft position under rejection sampling. Measured cost: 5.6-8.0 ms/call; the actual work is finding the few-hundred-token head whose cumulative mass reaches g_nuc. Replace the full qsort + linear scan with a max-heap partial select: - Floyd heapify g_pidx over V by descending g_pbuf prob (O(V), cache-friendly) - pop winners to the array's high end until cum >= g_nuc (k * O(log V)) - the remaining heap prefix IS the tail -> zero it, renormalize the head Winners land in g_pidx[out..V-1] in descending order, so s2 accumulates in the same order as before -> head is unchanged on tie-free shapes (ties were already unspecified under the unstable qsort). All four dist_build/dist_sample contract properties hold: g_pbuf stays id-indexed, g_pidx stays internal, the tail is fully zeroed, the head renormalizes to 1. No API change, no caller change, no new globals. c/tests/test_topp.c (new): drives the REAL dist_build via the include-glm.c pattern against an independent double-precision reimplementation of the OLD algorithm. 123 cases: 6 sizes (1..1519) x 5 shapes (uniform/peaked/geometric/ plateau/sharptail) x 4 nuc values, plus the g_nuc>=1 guard-off paths and V=1. Tie-free shapes compare head values within 1e-6 rel (float vs double renorm noise); tie shapes compare value-multisets. No scratch files -> builds clean on Windows MinGW without the unmerged mkdtemp shim (#352).
This commit is contained in:
+4
-1
@@ -166,7 +166,7 @@ else
|
||||
PYTHON ?= python3
|
||||
endif
|
||||
CUDA_OBJ =
|
||||
TEST_BINS = tests/test_json$(EXE) tests/test_st$(EXE) tests/test_tier$(EXE) tests/test_grammar$(EXE) tests/test_schema_gbnf$(EXE) tests/test_decode_batch$(EXE) tests/test_idot$(EXE) tests/test_i4_grouped$(EXE) tests/test_stops$(EXE) tests/test_kv_alloc$(EXE) tests/test_i4_acc512$(EXE) tests/test_compat_direct$(EXE)
|
||||
TEST_BINS = tests/test_json$(EXE) tests/test_st$(EXE) tests/test_tier$(EXE) tests/test_grammar$(EXE) tests/test_schema_gbnf$(EXE) tests/test_decode_batch$(EXE) tests/test_idot$(EXE) tests/test_i4_grouped$(EXE) tests/test_stops$(EXE) tests/test_topp$(EXE) tests/test_kv_alloc$(EXE) tests/test_i4_acc512$(EXE) tests/test_compat_direct$(EXE)
|
||||
ifneq (,$(LINUX))
|
||||
TEST_BINS += tests/test_uring$(EXE)
|
||||
endif
|
||||
@@ -317,6 +317,9 @@ tests/test_i4_grouped$(EXE): tests/test_i4_grouped.c glm.c st.h uring.h json.h t
|
||||
tests/test_stops$(EXE): tests/test_stops.c glm.c st.h uring.h json.h tok.h tok_unicode.h compat.h grammar.h tier.h
|
||||
$(CC) $(CFLAGS) $< -o $@ $(LDFLAGS)
|
||||
|
||||
tests/test_topp$(EXE): tests/test_topp.c glm.c st.h uring.h json.h tok.h tok_unicode.h compat.h grammar.h tier.h
|
||||
$(CC) $(CFLAGS) $< -o $@ $(LDFLAGS)
|
||||
|
||||
tests/test_kv_alloc$(EXE): tests/test_kv_alloc.c glm.c st.h json.h tok.h tok_unicode.h compat.h grammar.h tier.h
|
||||
$(CC) $(CFLAGS) $< -o $@ $(LDFLAGS)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user