| import test from 'node:test'; |
| import assert from 'node:assert/strict'; |
| import { nucleusProcessor } from '../src/sampling.mjs'; |
|
|
| const filter = (values, p) => { |
| const data = Float32Array.from(values); |
| nucleusProcessor(p)([], { dims: [1, data.length], data }); |
| return [...data].map(Number.isFinite); |
| }; |
| test('nucleus keeps enough probability, includes the boundary token, and retains one token', () => { |
| const scores = [.6, .25, .1, .05].map(Math.log); |
| assert.deepEqual(filter(scores, .8), [true, true, false, false]); |
| assert.deepEqual(filter(scores, .99), [true, true, true, true]); |
| assert.deepEqual(filter(scores, .01), [true, false, false, false]); |
| assert.deepEqual(filter(scores, 1), [true, true, true, true]); |
| assert.equal(filter([0, 0, 0, 0], .3).filter(Boolean).length, 2); |
| assert.equal(filter([-Infinity, -1000, 0], .95).filter(Boolean).length, 1); |
| assert.throws(() => filter([NaN, 0], .95), /invalid/); |
| assert.throws(() => nucleusProcessor(0)); |
| }); |
| test('tail optimization agrees with a full-sort reference over varied distributions', () => { |
| let seed = 42; |
| const random = () => ((seed = Math.imul(seed, 1664525) + 1013904223 >>> 0) / 2 ** 32); |
| for (const size of [2, 10, 1000, 130560]) for (const p of [.1, .5, .95, .999]) { |
| const scores = Float32Array.from({ length: size }, () => random() * 60 - 30); |
| const order = Array.from(scores, (value, i) => ({ value, i })).sort((a, b) => a.value - b.value || a.i - b.i); |
| const max = order.at(-1).value; |
| const total = order.reduce((sum, item) => sum + Math.exp(item.value - max), 0); |
| const expected = new Array(size).fill(true); |
| let cumulative = 0; |
| for (const item of order.slice(0, -1)) { |
| cumulative += Math.exp(item.value - max); |
| if (cumulative <= (1 - p) * total) expected[item.i] = false; |
| } |
| assert.deepEqual(filter(scores, p), expected, `size=${size}, p=${p}`); |
| } |
| }); |
|
|