File size: 1,902 Bytes
39371ea
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
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}`);
  }
});