webbrain-one's picture
Tiny XS v3.1: independently validated 32K WebGPU runtime
cbc5ce9 verified
Raw History Blame Contribute Delete
6.53 kB
import * as ort from '/runtime/vendor/ort.webgpu.mjs';
import {AutoTokenizer, env} from '/runtime/vendor/transformers.web.js';
let session, tokenizer, abi, gpuInfo, eos;
const MAX_CONTEXT_TOKENS = 32768;
const MAX_OUTPUT_TOKENS = 2048;
const PREFILL_CHUNK_TOKENS = 512;
function emptyCache() {
return abi.inputs.slice(2).map(() => new ort.Tensor('float16', new Uint16Array(0), [1, abi.numKvHeads, 0, abi.headDim]));
}
function dispose(tensors) { for (const t of tensors) if (t.location === 'gpu-buffer') t.dispose(); }
function argmax(values) { let best = 0; for (let i=1;i<values.length;i++) if(values[i]>values[best]) best=i; return best; }
window.initializeSpark = async () => {
console.log('spark-stage: adapter');
const adapter = await navigator.gpu.requestAdapter({powerPreference: 'high-performance'});
if (!adapter || !adapter.features.has('shader-f16')) throw Error('Hardware WebGPU/shader-f16 required');
const info = adapter.info;
if (info.vendor !== 'nvidia' || info.architecture !== 'blackwell') throw Error('Wrong WebGPU adapter');
gpuInfo = {vendor: info.vendor, architecture: info.architecture, description: info.description};
ort.env.webgpu.adapter = adapter;
ort.env.wasm.numThreads = 1;
ort.env.wasm.wasmPaths = '/runtime/vendor/';
env.allowRemoteModels = false;
env.allowLocalModels = true;
env.localModelPath = '/';
console.log('spark-stage: tokenizer');
tokenizer = await AutoTokenizer.from_pretrained('model', {local_files_only: true});
// Transformers.js 4.2 does not load the separate Jinja file automatically.
const templateResponse = await fetch('/model/chat_template.jinja');
if (!templateResponse.ok) throw Error('Missing native Spark chat template');
tokenizer.chat_template = (await templateResponse.text()).replace(/\r\n/g, '\n');
abi = await (await fetch('/abi')).json();
const config = await (await fetch('/model/generation_config.json')).json();
eos = new Set(Array.isArray(config.eos_token_id) ? config.eos_token_id : [config.eos_token_id]);
// Let ORT's streaming large-file loader bypass Chrome's 2GB arrayBuffer cap.
console.log('spark-stage: creating session');
session = await ort.InferenceSession.create('/model/onnx/model_fp16.onnx', {
executionProviders: ['webgpu'], graphOptimizationLevel: 'disabled',
preferredOutputLocation: 'gpu-buffer',
externalData: [{path: 'model_fp16.onnx_data', data: '/model/onnx/model_fp16.onnx_data'}],
});
return {gpuInfo, runtime: ort.env.versions, abi};
};
async function step(ids, past, seen) {
if (!ids.length || seen + ids.length > MAX_CONTEXT_TOKENS) throw Error('Context limit exceeded');
const feeds = {
input_ids: new ort.Tensor('int64', BigInt64Array.from(ids, BigInt), [1, ids.length]),
position_ids: new ort.Tensor('int64', BigInt64Array.from(ids.map((_,i)=>seen+i), BigInt), [1,ids.length]),
};
abi.inputs.slice(2).forEach((key,i)=>feeds[key]=past[i]);
let result, logits, present;
try {
result = await session.run(feeds);
logits = await result.logits.getData();
if (!logits.every(Number.isFinite)) throw Error('Nonfinite WebGPU logits');
present = abi.outputs.slice(1).map(key=>result[key]);
if (!present.every(t=>t.dims[2]===seen+ids.length)) throw Error('Cache-length mismatch');
} catch (error) {
if (result) dispose(Object.values(result));
throw error;
}
result.logits.dispose();
dispose(past);
return {logits, present};
}
// Keep every token. Small, ordered prefill steps bound transient attention memory;
// each returned cache becomes the next step's input without a context reset.
async function prefill(ids) {
if (!Array.isArray(ids) || ids.length < 1 || ids.length > MAX_CONTEXT_TOKENS) {
throw Error('Prompt must contain 1..32768 tokens');
}
let past = emptyCache(), seen = 0, logits;
try {
for (let offset = 0; offset < ids.length; offset += PREFILL_CHUNK_TOKENS) {
const chunk = ids.slice(offset, offset + PREFILL_CHUNK_TOKENS);
const result = await step(chunk, past, seen);
past = result.present;
seen += chunk.length;
logits = result.logits;
}
if (seen !== ids.length || past.some(t => t.dims[2] !== seen)) {
throw Error('Prefill lost prompt tokens');
}
return {logits, past, seen};
} catch (error) {
dispose(past);
throw error;
}
}
window.sparkNumerics = async (fixture) => {
if (!Array.isArray(fixture.teacherTokens) || fixture.teacherTokens.length < 3 ||
fixture.inputIds.length + 3 > MAX_CONTEXT_TOKENS) throw Error('Invalid numerical fixture');
const initial = await prefill(fixture.inputIds);
let past = initial.past, seen = initial.seen;
const logits=[];
try {
logits.push(Array.from(initial.logits));
for(let n=0;n<3;n++) {
const result=await step([fixture.teacherTokens[n]],past,seen);
logits.push(Array.from(result.logits));
seen++; past=result.present;
}
return {logits,finite:true,gpuInfo,finalCacheTokens:seen};
} finally {
dispose(past);
}
};
window.sparkTemplate = ({messages,tools}) => {
const rendered=tokenizer.apply_chat_template(messages,{tools:tools||undefined,tokenize:false,
enable_thinking:false,add_generation_prompt:true});
return {text:rendered,ids:Array.from(tokenizer(rendered,{add_special_tokens:false}).input_ids.data,Number)};
};
window.sparkGenerate = async ({messages,tools,maxTokens=512}) => {
if (!Number.isInteger(maxTokens) || maxTokens < 1 || maxTokens > MAX_OUTPUT_TOKENS) {
throw Error('maxTokens must be 1..2048');
}
const prompt=window.sparkTemplate({messages,tools});
if(prompt.ids.length+maxTokens>MAX_CONTEXT_TOKENS) {
throw Error('Input plus requested output exceeds 32768 tokens; no truncation applied');
}
let past, seen=0, output=[], reason='length';
const started=performance.now();
try {
const initial = await prefill(prompt.ids);
past = initial.past;
seen = initial.seen;
let logits = initial.logits;
for(let n=0;n<maxTokens;n++) {
const token=argmax(logits);
output.push(token);
if(eos.has(token)){reason='stop';break;}
if (n + 1 < maxTokens) {
const result=await step([token],past,seen);
past=result.present;seen++;
logits=result.logits;
}
}
return {content:tokenizer.decode(output,{skip_special_tokens:true}),promptTokens:prompt.ids.length,
completionTokens:output.length,finishReason:reason,seconds:(performance.now()-started)/1000};
} finally {if (past) dispose(past);}
};