| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| import { RandomStream } from '../utils/rng'; |
| import { endsSentence } from '../utils/sentences'; |
| import type { CandidateInfo, StepTrace } from '../watermark/types'; |
|
|
| |
| export interface LoopModel { |
| |
| |
| |
| |
| forward( |
| inputIds: number[], |
| fullSeqLen: number, |
| past: unknown, |
| ): Promise<{ logits: Float32Array; past: unknown }>; |
| |
| disposePast(past: unknown): void; |
| vocabSize: number; |
| eosTokenIds: number[]; |
| decode(ids: number[]): string; |
| } |
|
|
| export interface GenerationHooks { |
| |
| transformLogits?: (stepIndex: number, contextIds: number[], logits: Float32Array) => bigint | void; |
| |
| sampleOverride?: ( |
| stepIndex: number, |
| contextIds: number[], |
| probs: Float32Array, |
| ) => { tokenId: number; keyId?: 1 | 2; r?: number; gumbelScore?: number }; |
| |
| onSentenceEnd?: ( |
| sentenceText: string, |
| sentenceIndex: number, |
| attempt: number, |
| ) => Promise<{ accept: boolean }>; |
| maxSentenceTrials?: number; |
| } |
|
|
| export interface LoopConfig { |
| temperature: number; |
| topP: number; |
| maxNewTokens: number; |
| |
| baseSeed: number; |
| topKTrace: number; |
| |
| |
| |
| |
| |
| |
| |
| |
| retryTemperatureStep?: number; |
| |
| retryTemperatureMax?: number; |
| |
| |
| |
| |
| |
| |
| onToken?: (step: StepTrace) => void; |
| } |
|
|
| export interface LoopResult { |
| tokenIds: number[]; |
| text: string; |
| steps: StepTrace[]; |
| retries: number; |
| aborted: boolean; |
| } |
|
|
| |
| export function softmaxT(logits: Float32Array, temperature: number): Float32Array { |
| const out = new Float32Array(logits.length); |
| const t = Math.max(temperature, 1e-4); |
| let max = -Infinity; |
| for (let i = 0; i < logits.length; i++) if (logits[i] > max) max = logits[i]; |
| let sum = 0; |
| for (let i = 0; i < logits.length; i++) { |
| const e = Math.exp((logits[i] - max) / t); |
| out[i] = e; |
| sum += e; |
| } |
| for (let i = 0; i < out.length; i++) out[i] /= sum; |
| return out; |
| } |
|
|
| |
| export function applyTopP(probs: Float32Array, topP: number): void { |
| if (topP >= 1) return; |
| const idx = Array.from(probs.keys()); |
| idx.sort((a, b) => probs[b] - probs[a]); |
| let cum = 0; |
| let cut = idx.length; |
| for (let i = 0; i < idx.length; i++) { |
| cum += probs[idx[i]]; |
| if (cum >= topP) { |
| cut = i + 1; |
| break; |
| } |
| } |
| const keep = new Set(idx.slice(0, cut)); |
| let sum = 0; |
| for (let i = 0; i < probs.length; i++) { |
| if (!keep.has(i)) probs[i] = 0; |
| else sum += probs[i]; |
| } |
| if (sum > 0) for (let i = 0; i < probs.length; i++) probs[i] /= sum; |
| } |
|
|
| |
| export function sampleFromProbs(probs: Float32Array, u: number): number { |
| let cum = 0; |
| for (let i = 0; i < probs.length; i++) { |
| cum += probs[i]; |
| if (u < cum) return i; |
| } |
| |
| for (let i = probs.length - 1; i >= 0; i--) if (probs[i] > 0) return i; |
| return 0; |
| } |
|
|
| |
| export function entropyOf(probs: Float32Array): number { |
| let h = 0; |
| for (let i = 0; i < probs.length; i++) { |
| const p = probs[i]; |
| if (p > 0) h -= p * Math.log(p); |
| } |
| return h; |
| } |
|
|
| function topKCandidates( |
| preLogits: Float32Array, |
| preProbs: Float32Array, |
| postProbs: Float32Array | null, |
| k: number, |
| decode: (ids: number[]) => string, |
| ): CandidateInfo[] { |
| const idx = Array.from(preProbs.keys()); |
| idx.sort((a, b) => preProbs[b] - preProbs[a]); |
| return idx.slice(0, k).map((v) => ({ |
| tokenId: v, |
| tokenText: decode([v]), |
| logit: preLogits[v], |
| prob: preProbs[v], |
| probAfter: postProbs ? postProbs[v] : undefined, |
| })); |
| } |
|
|
| export async function runGenerationLoop( |
| model: LoopModel, |
| promptIds: number[], |
| cfg: LoopConfig, |
| hooks: GenerationHooks = {}, |
| ): Promise<LoopResult> { |
| const steps: StepTrace[] = []; |
| const generated: number[] = []; |
| const stream = new RandomStream(BigInt(cfg.baseSeed)); |
| let retryStreamSalt = 1; |
| let retries = 0; |
|
|
| let seq = [...promptIds]; |
| let past: unknown = null; |
| let pending = [...promptIds]; |
| let sentenceStartLen = 0; |
| let sentenceIndex = 0; |
| let attempt = 1; |
| let aborted = false; |
| |
| |
| |
| |
| |
| |
| |
| |
| const promptText = model.decode(promptIds); |
| let decodedSoFar = ''; |
|
|
| |
| function continuationOf(ids: number[]): string { |
| const whole = model.decode([...promptIds, ...ids]); |
| return whole.startsWith(promptText) ? whole.slice(promptText.length) : model.decode(ids); |
| } |
|
|
| |
| |
| |
| |
| const settled = (s: string) => s.replace(/�+$/, ''); |
|
|
| function textAddedBy(chosen: number): string { |
| const before = settled(decodedSoFar); |
| decodedSoFar = continuationOf([...generated, chosen]); |
| const now = settled(decodedSoFar); |
| |
| |
| return now.startsWith(before) ? now.slice(before.length) : now; |
| } |
|
|
| const maxTrials = hooks.maxSentenceTrials ?? 12; |
|
|
| try { |
| while (generated.length < cfg.maxNewTokens) { |
| const { logits, past: newPast } = await model.forward(pending, seq.length, past); |
| past = newPast; |
| pending = []; |
|
|
| |
| |
| const temperature = |
| cfg.temperature * |
| Math.min( |
| cfg.retryTemperatureMax ?? 1.5, |
| 1 + (cfg.retryTemperatureStep ?? 0) * (attempt - 1), |
| ); |
|
|
| |
| const preLogits = logits.slice(); |
| const preProbs = softmaxT(preLogits, temperature); |
| const entropy = entropyOf(preProbs); |
|
|
| const contextIds = seq; |
| const seed = hooks.transformLogits?.(generated.length, contextIds, logits); |
|
|
| let probs = softmaxT(logits, temperature); |
| applyTopP(probs, cfg.topP); |
|
|
| let chosen: number; |
| let keyId: 1 | 2 | undefined; |
| let r: number | undefined; |
| let gumbelScore: number | undefined; |
| if (hooks.sampleOverride) { |
| const pick = hooks.sampleOverride(generated.length, contextIds, probs); |
| chosen = pick.tokenId; |
| keyId = pick.keyId; |
| r = pick.r; |
| gumbelScore = pick.gumbelScore; |
| } else { |
| chosen = sampleFromProbs(probs, stream.next()); |
| } |
|
|
| const isEosToken = model.eosTokenIds.includes(chosen); |
| const postProbs = seed !== undefined ? probs : null; |
| const stepTrace: StepTrace = { |
| index: generated.length, |
| chosenTokenId: chosen, |
| chosenTokenText: isEosToken ? '' : textAddedBy(chosen), |
| seed: seed !== undefined && seed !== null ? String(seed) : undefined, |
| keyId, |
| topCandidates: topKCandidates(preLogits, preProbs, postProbs, cfg.topKTrace, (ids) => |
| model.decode(ids), |
| ), |
| entropy, |
| }; |
| if (r !== undefined) { |
| stepTrace.topCandidates.forEach((c) => { |
| if (c.tokenId === chosen) { |
| c.r = r; |
| c.gumbelScore = gumbelScore; |
| } |
| }); |
| } |
|
|
| const isEos = isEosToken; |
| if (!isEos) { |
| seq = [...seq, chosen]; |
| pending = [chosen]; |
| generated.push(chosen); |
| steps.push(stepTrace); |
| cfg.onToken?.(stepTrace); |
| } |
|
|
| |
| if (hooks.onSentenceEnd) { |
| const sentText = model.decode(generated.slice(sentenceStartLen)); |
| const boundary = isEos || generated.length >= cfg.maxNewTokens || endsSentence(sentText); |
| if (boundary && generated.length > sentenceStartLen) { |
| const { accept } = await hooks.onSentenceEnd(sentText.trim(), sentenceIndex, attempt); |
| if (accept || attempt >= maxTrials) { |
| sentenceStartLen = generated.length; |
| sentenceIndex++; |
| attempt = 1; |
| } else { |
| |
| retries++; |
| attempt++; |
| const keep = generated.slice(0, sentenceStartLen); |
| const removedSteps = generated.length - sentenceStartLen; |
| generated.length = sentenceStartLen; |
| steps.length = steps.length - removedSteps; |
| decodedSoFar = continuationOf(generated); |
| seq = [...promptIds, ...keep]; |
| model.disposePast(past); |
| past = null; |
| pending = [...seq]; |
| |
| for (let i = 0; i < retryStreamSalt; i++) stream.next(); |
| retryStreamSalt++; |
| if (isEos) continue; |
| } |
| } |
| } |
|
|
| if (isEos) break; |
| } |
| } catch (e) { |
| aborted = true; |
| throw e; |
| } finally { |
| model.disposePast(past); |
| } |
|
|
| return { |
| tokenIds: generated, |
| |
| |
| text: continuationOf(generated), |
| steps, |
| retries, |
| aborted, |
| }; |
| } |
|
|