MiniCPM5-2B-WebGPU-Pi-HTTP / app /tests /sampling.test.mjs
Mike0021's picture
Enable browser HTTP commands with CORS-aware errors and bounded requests
39371ea verified
Raw
History Blame Contribute Delete
1.9 kB
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}`);
}
});