import * as ort from '/runtime/vendor/ort.webgpu.mjs'; import {AutoTokenizer, env} from '/runtime/vendor/transformers.web.js'; let session, tokenizer, abi, gpuInfo, eos; 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;ivalues[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) { 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]); const result = await session.run(feeds); const logits = await result.logits.getData(); if (![...logits].every(Number.isFinite)) throw Error('Nonfinite WebGPU logits'); const present = abi.outputs.slice(1).map(key=>result[key]); if (!present.every(t=>t.dims[2]===seen+ids.length)) throw Error('Cache-length mismatch'); result.logits.dispose(); dispose(past); return {logits, present}; } window.sparkNumerics = async (fixture) => { let ids=fixture.inputIds, past=emptyCache(), seen=0; const logits=[]; for(let n=0;n<4;n++) { const result=await step(ids,past,seen); logits.push(Array.from(result.logits)); seen+=ids.length; ids=[fixture.teacherTokens[n]]; past=result.present; } dispose(past); return {logits,finite:true,gpuInfo}; }; 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}) => { const prompt=window.sparkTemplate({messages,tools}); if(prompt.ids.length+maxTokens>4096) throw Error('This deployed WebGPU endpoint supports at most 4096 input+output tokens'); let ids=prompt.ids,past=emptyCache(),seen=0,output=[],reason='length'; const started=performance.now(); try { for(let n=0;n