| |
| |
| |
| export function nucleusProcessor(topP = 0.95) { |
| if (!(topP > 0 && topP <= 1)) throw Error('topP must be in (0, 1].'); |
| let sorted; |
| return (_inputIds, logits) => { |
| if (topP === 1) return logits; |
| const size = logits.dims.at(-1); |
| sorted ??= new Float32Array(size); |
| if (sorted.length !== size) sorted = new Float32Array(size); |
| for (let offset = 0; offset < logits.data.length; offset += size) { |
| const scores = logits.data.subarray(offset, offset + size); |
| let max = -Infinity; |
| for (let i = 0; i < size; i++) max = Math.max(max, scores[i]); |
| if (!Number.isFinite(max)) throw Error('The model produced invalid sampling scores.'); |
| let total = 0; |
| for (let i = 0; i < size; i++) total += Math.exp(scores[i] - max); |
| const tailMass = (1 - topP) * total; |
| |
| |
| let floor = max - 8, count, mass; |
| do { |
| count = 0; mass = 0; |
| for (let i = 0; i < size; i++) { |
| if (scores[i] >= floor) sorted[count++] = scores[i]; |
| else mass += Math.exp(scores[i] - max); |
| } |
| if (mass <= tailMass) break; |
| floor -= 8; |
| } while (true); |
| sorted.subarray(0, count).sort(); |
| let removed = 0; |
| |
| while (removed < count - 1) { |
| const next = mass + Math.exp(sorted[removed] - max); |
| if (next > tailMass) break; |
| mass = next; removed++; |
| } |
| const cutoff = removed ? sorted[removed - 1] : floor; |
| let tied = 0; |
| for (let i = removed - 1; i >= 0 && sorted[i] === cutoff; i--) tied++; |
| for (let i = 0; i < size; i++) { |
| if (scores[i] < cutoff || (scores[i] === cutoff && tied-- > 0)) scores[i] = -Infinity; |
| } |
| } |
| return logits; |
| }; |
| } |
|
|