From f225c7c2a6cc59d0b20aafda6eee44799eeb95da Mon Sep 17 00:00:00 2001 From: w1n5t0n Date: Fri, 3 Apr 2026 17:11:49 +0100 Subject: [PATCH] feat(wasm): add batch inference, extended training, pin mask, eval loss, and layer stats bindings Five new C functions for the WASM module: - nisps_mlp_infer_batch: N-point batch inference in a single call - nisps_mlp_train_ex: training with per-iteration loss history output - nisps_mlp_move_weights_ex: moveWeights with output pin mask to skip pinned nodes - nisps_mlp_eval_loss: compute MSE loss without updating weights - nisps_mlp_get_layer_stats: per-layer weight magnitude, dead, and saturation stats --- playground/wasm/build.sh | 5 + playground/wasm/nisps.js | 2 +- playground/wasm/nisps.wasm | Bin 33776 -> 36699 bytes playground/wasm/nisps_bindings.cpp | 182 +++++++++++++++++++++++++++++ 4 files changed, 188 insertions(+), 1 deletion(-) diff --git a/playground/wasm/build.sh b/playground/wasm/build.sh index b0724ec..4e1f3e2 100755 --- a/playground/wasm/build.sh +++ b/playground/wasm/build.sh @@ -22,6 +22,11 @@ emcc "$SCRIPT_DIR/nisps_bindings.cpp" \ "_nisps_mlp_train", "_nisps_mlp_draw_weights_spread", "_nisps_mlp_move_weights_spread", + "_nisps_mlp_infer_batch", + "_nisps_mlp_train_ex", + "_nisps_mlp_move_weights_ex", + "_nisps_mlp_eval_loss", + "_nisps_mlp_get_layer_stats", "_nisps_alloc", "_nisps_free", "_nisps_alloc_int", diff --git a/playground/wasm/nisps.js b/playground/wasm/nisps.js index 2795a46..52aad07 100644 --- a/playground/wasm/nisps.js +++ b/playground/wasm/nisps.js @@ -1,2 +1,2 @@ -async function NispsModule(moduleArg={}){var moduleRtn;var Module=moduleArg;var ENVIRONMENT_IS_WEB=!!globalThis.window;var ENVIRONMENT_IS_WORKER=!!globalThis.WorkerGlobalScope;var ENVIRONMENT_IS_NODE=globalThis.process?.versions?.node&&globalThis.process?.type!="renderer";var arguments_=[];var thisProgram="./this.program";var _scriptName=import.meta.url;var scriptDirectory="";function locateFile(path){if(Module["locateFile"]){return Module["locateFile"](path,scriptDirectory)}return scriptDirectory+path}var readAsync,readBinary;if(ENVIRONMENT_IS_WEB||ENVIRONMENT_IS_WORKER){try{scriptDirectory=new URL(".",_scriptName).href}catch{}{if(ENVIRONMENT_IS_WORKER){readBinary=url=>{var xhr=new XMLHttpRequest;xhr.open("GET",url,false);xhr.responseType="arraybuffer";xhr.send(null);return new Uint8Array(xhr.response)}}readAsync=async url=>{var response=await fetch(url,{credentials:"same-origin"});if(response.ok){return response.arrayBuffer()}throw new Error(response.status+" : "+response.url)}}}else{}var out=console.log.bind(console);var err=console.error.bind(console);var wasmBinary;var ABORT=false;var readyPromiseResolve,readyPromiseReject;var HEAP8,HEAPU8,HEAP16,HEAPU16,HEAP32,HEAPU32,HEAPF32,HEAPF64;var HEAP64,HEAPU64;var runtimeInitialized=false;function updateMemoryViews(){var b=wasmMemory.buffer;HEAP8=new Int8Array(b);HEAP16=new Int16Array(b);HEAPU8=new Uint8Array(b);HEAPU16=new Uint16Array(b);Module["HEAP32"]=HEAP32=new Int32Array(b);HEAPU32=new Uint32Array(b);Module["HEAPF32"]=HEAPF32=new Float32Array(b);HEAPF64=new Float64Array(b);HEAP64=new BigInt64Array(b);HEAPU64=new BigUint64Array(b)}function preRun(){if(Module["preRun"]){if(typeof Module["preRun"]=="function")Module["preRun"]=[Module["preRun"]];while(Module["preRun"].length){addOnPreRun(Module["preRun"].shift())}}callRuntimeCallbacks(onPreRuns)}function initRuntime(){runtimeInitialized=true;wasmExports["__wasm_call_ctors"]()}function postRun(){if(Module["postRun"]){if(typeof Module["postRun"]=="function")Module["postRun"]=[Module["postRun"]];while(Module["postRun"].length){addOnPostRun(Module["postRun"].shift())}}callRuntimeCallbacks(onPostRuns)}function abort(what){Module["onAbort"]?.(what);what="Aborted("+what+")";err(what);ABORT=true;what+=". Build with -sASSERTIONS for more info.";var e=new WebAssembly.RuntimeError(what);readyPromiseReject?.(e);throw e}var wasmBinaryFile;function findWasmBinary(){if(Module["locateFile"]){return locateFile("nisps.wasm")}return new URL("nisps.wasm",import.meta.url).href}function getBinarySync(file){if(file==wasmBinaryFile&&wasmBinary){return new Uint8Array(wasmBinary)}if(readBinary){return readBinary(file)}throw"both async and sync fetching of the wasm failed"}async function getWasmBinary(binaryFile){if(!wasmBinary){try{var response=await readAsync(binaryFile);return new Uint8Array(response)}catch{}}return getBinarySync(binaryFile)}async function instantiateArrayBuffer(binaryFile,imports){try{var binary=await getWasmBinary(binaryFile);var instance=await WebAssembly.instantiate(binary,imports);return instance}catch(reason){err(`failed to asynchronously prepare wasm: ${reason}`);abort(reason)}}async function instantiateAsync(binary,binaryFile,imports){if(!binary){try{var response=fetch(binaryFile,{credentials:"same-origin"});var instantiationResult=await WebAssembly.instantiateStreaming(response,imports);return instantiationResult}catch(reason){err(`wasm streaming compile failed: ${reason}`);err("falling back to ArrayBuffer instantiation")}}return instantiateArrayBuffer(binaryFile,imports)}function getWasmImports(){var imports={env:wasmImports,wasi_snapshot_preview1:wasmImports};return imports}async function createWasm(){function receiveInstance(instance,module){wasmExports=instance.exports;assignWasmExports(wasmExports);updateMemoryViews();return wasmExports}function receiveInstantiationResult(result){return receiveInstance(result["instance"])}var info=getWasmImports();if(Module["instantiateWasm"]){return new Promise((resolve,reject)=>{Module["instantiateWasm"](info,(inst,mod)=>{resolve(receiveInstance(inst,mod))})})}wasmBinaryFile??=findWasmBinary();var result=await instantiateAsync(wasmBinary,wasmBinaryFile,info);var exports=receiveInstantiationResult(result);return exports}class ExitStatus{name="ExitStatus";constructor(status){this.message=`Program terminated with exit(${status})`;this.status=status}}var callRuntimeCallbacks=callbacks=>{while(callbacks.length>0){callbacks.shift()(Module)}};var onPostRuns=[];var addOnPostRun=cb=>onPostRuns.push(cb);var onPreRuns=[];var addOnPreRun=cb=>onPreRuns.push(cb);var noExitRuntime=true;var stackRestore=val=>__emscripten_stack_restore(val);var stackSave=()=>_emscripten_stack_get_current();var UTF8Decoder=globalThis.TextDecoder&&new TextDecoder;var findStringEnd=(heapOrArray,idx,maxBytesToRead,ignoreNul)=>{var maxIdx=idx+maxBytesToRead;if(ignoreNul)return maxIdx;while(heapOrArray[idx]&&!(idx>=maxIdx))++idx;return idx};var UTF8ArrayToString=(heapOrArray,idx=0,maxBytesToRead,ignoreNul)=>{var endPtr=findStringEnd(heapOrArray,idx,maxBytesToRead,ignoreNul);if(endPtr-idx>16&&heapOrArray.buffer&&UTF8Decoder){return UTF8Decoder.decode(heapOrArray.subarray(idx,endPtr))}var str="";while(idx>10,56320|ch&1023)}}return str};var UTF8ToString=(ptr,maxBytesToRead,ignoreNul)=>ptr?UTF8ArrayToString(HEAPU8,ptr,maxBytesToRead,ignoreNul):"";var ___assert_fail=(condition,filename,line,func)=>abort(`Assertion failed: ${UTF8ToString(condition)}, at: `+[filename?UTF8ToString(filename):"unknown filename",line,func?UTF8ToString(func):"unknown function"]);class ExceptionInfo{constructor(excPtr){this.excPtr=excPtr;this.ptr=excPtr-24}set_type(type){HEAPU32[this.ptr+4>>2]=type}get_type(){return HEAPU32[this.ptr+4>>2]}set_destructor(destructor){HEAPU32[this.ptr+8>>2]=destructor}get_destructor(){return HEAPU32[this.ptr+8>>2]}set_caught(caught){caught=caught?1:0;HEAP8[this.ptr+12]=caught}get_caught(){return HEAP8[this.ptr+12]!=0}set_rethrown(rethrown){rethrown=rethrown?1:0;HEAP8[this.ptr+13]=rethrown}get_rethrown(){return HEAP8[this.ptr+13]!=0}init(type,destructor){this.set_adjusted_ptr(0);this.set_type(type);this.set_destructor(destructor)}set_adjusted_ptr(adjustedPtr){HEAPU32[this.ptr+16>>2]=adjustedPtr}get_adjusted_ptr(){return HEAPU32[this.ptr+16>>2]}}var exceptionLast=0;var uncaughtExceptionCount=0;var ___cxa_throw=(ptr,type,destructor)=>{var info=new ExceptionInfo(ptr);info.init(type,destructor);exceptionLast=ptr;uncaughtExceptionCount++;throw exceptionLast};var __abort_js=()=>abort("");var getHeapMax=()=>2147483648;var alignMemory=(size,alignment)=>Math.ceil(size/alignment)*alignment;var growMemory=size=>{var oldHeapSize=wasmMemory.buffer.byteLength;var pages=(size-oldHeapSize+65535)/65536|0;try{wasmMemory.grow(pages);updateMemoryViews();return 1}catch(e){}};var _emscripten_resize_heap=requestedSize=>{var oldSize=HEAPU8.length;requestedSize>>>=0;var maxHeapSize=getHeapMax();if(requestedSize>maxHeapSize){return false}for(var cutDown=1;cutDown<=4;cutDown*=2){var overGrownHeapSize=oldSize*(1+.2/cutDown);overGrownHeapSize=Math.min(overGrownHeapSize,requestedSize+100663296);var newSize=Math.min(maxHeapSize,alignMemory(Math.max(requestedSize,overGrownHeapSize),65536));var replacement=growMemory(newSize);if(replacement){return true}}return false};var initRandomFill=()=>view=>crypto.getRandomValues(view);var randomFill=view=>{(randomFill=initRandomFill())(view)};var _random_get=(buffer,size)=>{randomFill(HEAPU8.subarray(buffer,buffer+size));return 0};var getCFunc=ident=>{var func=Module["_"+ident];return func};var writeArrayToMemory=(array,buffer)=>{HEAP8.set(array,buffer)};var lengthBytesUTF8=str=>{var len=0;for(var i=0;i=55296&&c<=57343){len+=4;++i}else{len+=3}}return len};var stringToUTF8Array=(str,heap,outIdx,maxBytesToWrite)=>{if(!(maxBytesToWrite>0))return 0;var startIdx=outIdx;var endIdx=outIdx+maxBytesToWrite-1;for(var i=0;i=endIdx)break;heap[outIdx++]=u}else if(u<=2047){if(outIdx+1>=endIdx)break;heap[outIdx++]=192|u>>6;heap[outIdx++]=128|u&63}else if(u<=65535){if(outIdx+2>=endIdx)break;heap[outIdx++]=224|u>>12;heap[outIdx++]=128|u>>6&63;heap[outIdx++]=128|u&63}else{if(outIdx+3>=endIdx)break;heap[outIdx++]=240|u>>18;heap[outIdx++]=128|u>>12&63;heap[outIdx++]=128|u>>6&63;heap[outIdx++]=128|u&63;i++}}heap[outIdx]=0;return outIdx-startIdx};var stringToUTF8=(str,outPtr,maxBytesToWrite)=>stringToUTF8Array(str,HEAPU8,outPtr,maxBytesToWrite);var stackAlloc=sz=>__emscripten_stack_alloc(sz);var stringToUTF8OnStack=str=>{var size=lengthBytesUTF8(str)+1;var ret=stackAlloc(size);stringToUTF8(str,ret,size);return ret};var ccall=(ident,returnType,argTypes,args,opts)=>{var toC={string:str=>{var ret=0;if(str!==null&&str!==undefined&&str!==0){ret=stringToUTF8OnStack(str)}return ret},array:arr=>{var ret=stackAlloc(arr.length);writeArrayToMemory(arr,ret);return ret}};function convertReturnValue(ret){if(returnType==="string"){return UTF8ToString(ret)}if(returnType==="boolean")return Boolean(ret);return ret}var func=getCFunc(ident);var cArgs=[];var stack=0;if(args){for(var i=0;i{var numericArgs=!argTypes||argTypes.every(type=>type==="number"||type==="boolean");var numericRet=returnType!=="string";if(numericRet&&numericArgs&&!opts){return getCFunc(ident)}return(...args)=>ccall(ident,returnType,argTypes,args,opts)};{if(Module["noExitRuntime"])noExitRuntime=Module["noExitRuntime"];if(Module["print"])out=Module["print"];if(Module["printErr"])err=Module["printErr"];if(Module["wasmBinary"])wasmBinary=Module["wasmBinary"];if(Module["arguments"])arguments_=Module["arguments"];if(Module["thisProgram"])thisProgram=Module["thisProgram"];if(Module["preInit"]){if(typeof Module["preInit"]=="function")Module["preInit"]=[Module["preInit"]];while(Module["preInit"].length>0){Module["preInit"].shift()()}}}Module["cwrap"]=cwrap;var _nisps_mlp_create,_nisps_mlp_destroy,_nisps_mlp_weight_count,_nisps_mlp_get_weights,_nisps_mlp_set_weights,_nisps_mlp_inference,_nisps_mlp_train,_nisps_mlp_draw_weights_spread,_nisps_mlp_move_weights_spread,_nisps_alloc,_malloc,_nisps_free,_free,_nisps_alloc_int,_nisps_free_int,__emscripten_stack_restore,__emscripten_stack_alloc,_emscripten_stack_get_current,memory,__indirect_function_table,wasmMemory;function assignWasmExports(wasmExports){_nisps_mlp_create=Module["_nisps_mlp_create"]=wasmExports["nisps_mlp_create"];_nisps_mlp_destroy=Module["_nisps_mlp_destroy"]=wasmExports["nisps_mlp_destroy"];_nisps_mlp_weight_count=Module["_nisps_mlp_weight_count"]=wasmExports["nisps_mlp_weight_count"];_nisps_mlp_get_weights=Module["_nisps_mlp_get_weights"]=wasmExports["nisps_mlp_get_weights"];_nisps_mlp_set_weights=Module["_nisps_mlp_set_weights"]=wasmExports["nisps_mlp_set_weights"];_nisps_mlp_inference=Module["_nisps_mlp_inference"]=wasmExports["nisps_mlp_inference"];_nisps_mlp_train=Module["_nisps_mlp_train"]=wasmExports["nisps_mlp_train"];_nisps_mlp_draw_weights_spread=Module["_nisps_mlp_draw_weights_spread"]=wasmExports["nisps_mlp_draw_weights_spread"];_nisps_mlp_move_weights_spread=Module["_nisps_mlp_move_weights_spread"]=wasmExports["nisps_mlp_move_weights_spread"];_nisps_alloc=Module["_nisps_alloc"]=wasmExports["nisps_alloc"];_malloc=Module["_malloc"]=wasmExports["malloc"];_nisps_free=Module["_nisps_free"]=wasmExports["nisps_free"];_free=Module["_free"]=wasmExports["free"];_nisps_alloc_int=Module["_nisps_alloc_int"]=wasmExports["nisps_alloc_int"];_nisps_free_int=Module["_nisps_free_int"]=wasmExports["nisps_free_int"];__emscripten_stack_restore=wasmExports["_emscripten_stack_restore"];__emscripten_stack_alloc=wasmExports["_emscripten_stack_alloc"];_emscripten_stack_get_current=wasmExports["emscripten_stack_get_current"];memory=wasmMemory=wasmExports["memory"];__indirect_function_table=wasmExports["__indirect_function_table"]}var wasmImports={__assert_fail:___assert_fail,__cxa_throw:___cxa_throw,_abort_js:__abort_js,emscripten_resize_heap:_emscripten_resize_heap,random_get:_random_get};function run(){preRun();function doRun(){Module["calledRun"]=true;if(ABORT)return;initRuntime();readyPromiseResolve?.(Module);Module["onRuntimeInitialized"]?.();postRun()}if(Module["setStatus"]){Module["setStatus"]("Running...");setTimeout(()=>{setTimeout(()=>Module["setStatus"](""),1);doRun()},1)}else{doRun()}}var wasmExports;wasmExports=await (createWasm());run();if(runtimeInitialized){moduleRtn=Module}else{moduleRtn=new Promise((resolve,reject)=>{readyPromiseResolve=resolve;readyPromiseReject=reject})} +async function NispsModule(moduleArg={}){var moduleRtn;var Module=moduleArg;var ENVIRONMENT_IS_WEB=!!globalThis.window;var ENVIRONMENT_IS_WORKER=!!globalThis.WorkerGlobalScope;var ENVIRONMENT_IS_NODE=globalThis.process?.versions?.node&&globalThis.process?.type!="renderer";var arguments_=[];var thisProgram="./this.program";var _scriptName=import.meta.url;var scriptDirectory="";function locateFile(path){if(Module["locateFile"]){return Module["locateFile"](path,scriptDirectory)}return scriptDirectory+path}var readAsync,readBinary;if(ENVIRONMENT_IS_WEB||ENVIRONMENT_IS_WORKER){try{scriptDirectory=new URL(".",_scriptName).href}catch{}{if(ENVIRONMENT_IS_WORKER){readBinary=url=>{var xhr=new XMLHttpRequest;xhr.open("GET",url,false);xhr.responseType="arraybuffer";xhr.send(null);return new Uint8Array(xhr.response)}}readAsync=async url=>{var response=await fetch(url,{credentials:"same-origin"});if(response.ok){return response.arrayBuffer()}throw new Error(response.status+" : "+response.url)}}}else{}var out=console.log.bind(console);var err=console.error.bind(console);var wasmBinary;var ABORT=false;var readyPromiseResolve,readyPromiseReject;var HEAP8,HEAPU8,HEAP16,HEAPU16,HEAP32,HEAPU32,HEAPF32,HEAPF64;var HEAP64,HEAPU64;var runtimeInitialized=false;function updateMemoryViews(){var b=wasmMemory.buffer;HEAP8=new Int8Array(b);HEAP16=new Int16Array(b);HEAPU8=new Uint8Array(b);HEAPU16=new Uint16Array(b);Module["HEAP32"]=HEAP32=new Int32Array(b);HEAPU32=new Uint32Array(b);Module["HEAPF32"]=HEAPF32=new Float32Array(b);HEAPF64=new Float64Array(b);HEAP64=new BigInt64Array(b);HEAPU64=new BigUint64Array(b)}function preRun(){if(Module["preRun"]){if(typeof Module["preRun"]=="function")Module["preRun"]=[Module["preRun"]];while(Module["preRun"].length){addOnPreRun(Module["preRun"].shift())}}callRuntimeCallbacks(onPreRuns)}function initRuntime(){runtimeInitialized=true;wasmExports["__wasm_call_ctors"]()}function postRun(){if(Module["postRun"]){if(typeof Module["postRun"]=="function")Module["postRun"]=[Module["postRun"]];while(Module["postRun"].length){addOnPostRun(Module["postRun"].shift())}}callRuntimeCallbacks(onPostRuns)}function abort(what){Module["onAbort"]?.(what);what="Aborted("+what+")";err(what);ABORT=true;what+=". Build with -sASSERTIONS for more info.";var e=new WebAssembly.RuntimeError(what);readyPromiseReject?.(e);throw e}var wasmBinaryFile;function findWasmBinary(){if(Module["locateFile"]){return locateFile("nisps.wasm")}return new URL("nisps.wasm",import.meta.url).href}function getBinarySync(file){if(file==wasmBinaryFile&&wasmBinary){return new Uint8Array(wasmBinary)}if(readBinary){return readBinary(file)}throw"both async and sync fetching of the wasm failed"}async function getWasmBinary(binaryFile){if(!wasmBinary){try{var response=await readAsync(binaryFile);return new Uint8Array(response)}catch{}}return getBinarySync(binaryFile)}async function instantiateArrayBuffer(binaryFile,imports){try{var binary=await getWasmBinary(binaryFile);var instance=await WebAssembly.instantiate(binary,imports);return instance}catch(reason){err(`failed to asynchronously prepare wasm: ${reason}`);abort(reason)}}async function instantiateAsync(binary,binaryFile,imports){if(!binary){try{var response=fetch(binaryFile,{credentials:"same-origin"});var instantiationResult=await WebAssembly.instantiateStreaming(response,imports);return instantiationResult}catch(reason){err(`wasm streaming compile failed: ${reason}`);err("falling back to ArrayBuffer instantiation")}}return instantiateArrayBuffer(binaryFile,imports)}function getWasmImports(){var imports={env:wasmImports,wasi_snapshot_preview1:wasmImports};return imports}async function createWasm(){function receiveInstance(instance,module){wasmExports=instance.exports;assignWasmExports(wasmExports);updateMemoryViews();return wasmExports}function receiveInstantiationResult(result){return receiveInstance(result["instance"])}var info=getWasmImports();if(Module["instantiateWasm"]){return new Promise((resolve,reject)=>{Module["instantiateWasm"](info,(inst,mod)=>{resolve(receiveInstance(inst,mod))})})}wasmBinaryFile??=findWasmBinary();var result=await instantiateAsync(wasmBinary,wasmBinaryFile,info);var exports=receiveInstantiationResult(result);return exports}class ExitStatus{name="ExitStatus";constructor(status){this.message=`Program terminated with exit(${status})`;this.status=status}}var callRuntimeCallbacks=callbacks=>{while(callbacks.length>0){callbacks.shift()(Module)}};var onPostRuns=[];var addOnPostRun=cb=>onPostRuns.push(cb);var onPreRuns=[];var addOnPreRun=cb=>onPreRuns.push(cb);var noExitRuntime=true;var stackRestore=val=>__emscripten_stack_restore(val);var stackSave=()=>_emscripten_stack_get_current();var UTF8Decoder=globalThis.TextDecoder&&new TextDecoder;var findStringEnd=(heapOrArray,idx,maxBytesToRead,ignoreNul)=>{var maxIdx=idx+maxBytesToRead;if(ignoreNul)return maxIdx;while(heapOrArray[idx]&&!(idx>=maxIdx))++idx;return idx};var UTF8ArrayToString=(heapOrArray,idx=0,maxBytesToRead,ignoreNul)=>{var endPtr=findStringEnd(heapOrArray,idx,maxBytesToRead,ignoreNul);if(endPtr-idx>16&&heapOrArray.buffer&&UTF8Decoder){return UTF8Decoder.decode(heapOrArray.subarray(idx,endPtr))}var str="";while(idx>10,56320|ch&1023)}}return str};var UTF8ToString=(ptr,maxBytesToRead,ignoreNul)=>ptr?UTF8ArrayToString(HEAPU8,ptr,maxBytesToRead,ignoreNul):"";var ___assert_fail=(condition,filename,line,func)=>abort(`Assertion failed: ${UTF8ToString(condition)}, at: `+[filename?UTF8ToString(filename):"unknown filename",line,func?UTF8ToString(func):"unknown function"]);class ExceptionInfo{constructor(excPtr){this.excPtr=excPtr;this.ptr=excPtr-24}set_type(type){HEAPU32[this.ptr+4>>2]=type}get_type(){return HEAPU32[this.ptr+4>>2]}set_destructor(destructor){HEAPU32[this.ptr+8>>2]=destructor}get_destructor(){return HEAPU32[this.ptr+8>>2]}set_caught(caught){caught=caught?1:0;HEAP8[this.ptr+12]=caught}get_caught(){return HEAP8[this.ptr+12]!=0}set_rethrown(rethrown){rethrown=rethrown?1:0;HEAP8[this.ptr+13]=rethrown}get_rethrown(){return HEAP8[this.ptr+13]!=0}init(type,destructor){this.set_adjusted_ptr(0);this.set_type(type);this.set_destructor(destructor)}set_adjusted_ptr(adjustedPtr){HEAPU32[this.ptr+16>>2]=adjustedPtr}get_adjusted_ptr(){return HEAPU32[this.ptr+16>>2]}}var exceptionLast=0;var uncaughtExceptionCount=0;var ___cxa_throw=(ptr,type,destructor)=>{var info=new ExceptionInfo(ptr);info.init(type,destructor);exceptionLast=ptr;uncaughtExceptionCount++;throw exceptionLast};var __abort_js=()=>abort("");var getHeapMax=()=>2147483648;var alignMemory=(size,alignment)=>Math.ceil(size/alignment)*alignment;var growMemory=size=>{var oldHeapSize=wasmMemory.buffer.byteLength;var pages=(size-oldHeapSize+65535)/65536|0;try{wasmMemory.grow(pages);updateMemoryViews();return 1}catch(e){}};var _emscripten_resize_heap=requestedSize=>{var oldSize=HEAPU8.length;requestedSize>>>=0;var maxHeapSize=getHeapMax();if(requestedSize>maxHeapSize){return false}for(var cutDown=1;cutDown<=4;cutDown*=2){var overGrownHeapSize=oldSize*(1+.2/cutDown);overGrownHeapSize=Math.min(overGrownHeapSize,requestedSize+100663296);var newSize=Math.min(maxHeapSize,alignMemory(Math.max(requestedSize,overGrownHeapSize),65536));var replacement=growMemory(newSize);if(replacement){return true}}return false};var initRandomFill=()=>view=>crypto.getRandomValues(view);var randomFill=view=>{(randomFill=initRandomFill())(view)};var _random_get=(buffer,size)=>{randomFill(HEAPU8.subarray(buffer,buffer+size));return 0};var getCFunc=ident=>{var func=Module["_"+ident];return func};var writeArrayToMemory=(array,buffer)=>{HEAP8.set(array,buffer)};var lengthBytesUTF8=str=>{var len=0;for(var i=0;i=55296&&c<=57343){len+=4;++i}else{len+=3}}return len};var stringToUTF8Array=(str,heap,outIdx,maxBytesToWrite)=>{if(!(maxBytesToWrite>0))return 0;var startIdx=outIdx;var endIdx=outIdx+maxBytesToWrite-1;for(var i=0;i=endIdx)break;heap[outIdx++]=u}else if(u<=2047){if(outIdx+1>=endIdx)break;heap[outIdx++]=192|u>>6;heap[outIdx++]=128|u&63}else if(u<=65535){if(outIdx+2>=endIdx)break;heap[outIdx++]=224|u>>12;heap[outIdx++]=128|u>>6&63;heap[outIdx++]=128|u&63}else{if(outIdx+3>=endIdx)break;heap[outIdx++]=240|u>>18;heap[outIdx++]=128|u>>12&63;heap[outIdx++]=128|u>>6&63;heap[outIdx++]=128|u&63;i++}}heap[outIdx]=0;return outIdx-startIdx};var stringToUTF8=(str,outPtr,maxBytesToWrite)=>stringToUTF8Array(str,HEAPU8,outPtr,maxBytesToWrite);var stackAlloc=sz=>__emscripten_stack_alloc(sz);var stringToUTF8OnStack=str=>{var size=lengthBytesUTF8(str)+1;var ret=stackAlloc(size);stringToUTF8(str,ret,size);return ret};var ccall=(ident,returnType,argTypes,args,opts)=>{var toC={string:str=>{var ret=0;if(str!==null&&str!==undefined&&str!==0){ret=stringToUTF8OnStack(str)}return ret},array:arr=>{var ret=stackAlloc(arr.length);writeArrayToMemory(arr,ret);return ret}};function convertReturnValue(ret){if(returnType==="string"){return UTF8ToString(ret)}if(returnType==="boolean")return Boolean(ret);return ret}var func=getCFunc(ident);var cArgs=[];var stack=0;if(args){for(var i=0;i{var numericArgs=!argTypes||argTypes.every(type=>type==="number"||type==="boolean");var numericRet=returnType!=="string";if(numericRet&&numericArgs&&!opts){return getCFunc(ident)}return(...args)=>ccall(ident,returnType,argTypes,args,opts)};{if(Module["noExitRuntime"])noExitRuntime=Module["noExitRuntime"];if(Module["print"])out=Module["print"];if(Module["printErr"])err=Module["printErr"];if(Module["wasmBinary"])wasmBinary=Module["wasmBinary"];if(Module["arguments"])arguments_=Module["arguments"];if(Module["thisProgram"])thisProgram=Module["thisProgram"];if(Module["preInit"]){if(typeof Module["preInit"]=="function")Module["preInit"]=[Module["preInit"]];while(Module["preInit"].length>0){Module["preInit"].shift()()}}}Module["cwrap"]=cwrap;var _nisps_mlp_create,_nisps_mlp_destroy,_nisps_mlp_weight_count,_nisps_mlp_get_weights,_nisps_mlp_set_weights,_nisps_mlp_inference,_nisps_mlp_train,_nisps_mlp_draw_weights_spread,_nisps_mlp_move_weights_spread,_nisps_mlp_infer_batch,_nisps_mlp_train_ex,_nisps_mlp_move_weights_ex,_nisps_mlp_eval_loss,_nisps_mlp_get_layer_stats,_nisps_alloc,_malloc,_nisps_free,_free,_nisps_alloc_int,_nisps_free_int,__emscripten_stack_restore,__emscripten_stack_alloc,_emscripten_stack_get_current,memory,__indirect_function_table,wasmMemory;function assignWasmExports(wasmExports){_nisps_mlp_create=Module["_nisps_mlp_create"]=wasmExports["nisps_mlp_create"];_nisps_mlp_destroy=Module["_nisps_mlp_destroy"]=wasmExports["nisps_mlp_destroy"];_nisps_mlp_weight_count=Module["_nisps_mlp_weight_count"]=wasmExports["nisps_mlp_weight_count"];_nisps_mlp_get_weights=Module["_nisps_mlp_get_weights"]=wasmExports["nisps_mlp_get_weights"];_nisps_mlp_set_weights=Module["_nisps_mlp_set_weights"]=wasmExports["nisps_mlp_set_weights"];_nisps_mlp_inference=Module["_nisps_mlp_inference"]=wasmExports["nisps_mlp_inference"];_nisps_mlp_train=Module["_nisps_mlp_train"]=wasmExports["nisps_mlp_train"];_nisps_mlp_draw_weights_spread=Module["_nisps_mlp_draw_weights_spread"]=wasmExports["nisps_mlp_draw_weights_spread"];_nisps_mlp_move_weights_spread=Module["_nisps_mlp_move_weights_spread"]=wasmExports["nisps_mlp_move_weights_spread"];_nisps_mlp_infer_batch=Module["_nisps_mlp_infer_batch"]=wasmExports["nisps_mlp_infer_batch"];_nisps_mlp_train_ex=Module["_nisps_mlp_train_ex"]=wasmExports["nisps_mlp_train_ex"];_nisps_mlp_move_weights_ex=Module["_nisps_mlp_move_weights_ex"]=wasmExports["nisps_mlp_move_weights_ex"];_nisps_mlp_eval_loss=Module["_nisps_mlp_eval_loss"]=wasmExports["nisps_mlp_eval_loss"];_nisps_mlp_get_layer_stats=Module["_nisps_mlp_get_layer_stats"]=wasmExports["nisps_mlp_get_layer_stats"];_nisps_alloc=Module["_nisps_alloc"]=wasmExports["nisps_alloc"];_malloc=Module["_malloc"]=wasmExports["malloc"];_nisps_free=Module["_nisps_free"]=wasmExports["nisps_free"];_free=Module["_free"]=wasmExports["free"];_nisps_alloc_int=Module["_nisps_alloc_int"]=wasmExports["nisps_alloc_int"];_nisps_free_int=Module["_nisps_free_int"]=wasmExports["nisps_free_int"];__emscripten_stack_restore=wasmExports["_emscripten_stack_restore"];__emscripten_stack_alloc=wasmExports["_emscripten_stack_alloc"];_emscripten_stack_get_current=wasmExports["emscripten_stack_get_current"];memory=wasmMemory=wasmExports["memory"];__indirect_function_table=wasmExports["__indirect_function_table"]}var wasmImports={__assert_fail:___assert_fail,__cxa_throw:___cxa_throw,_abort_js:__abort_js,emscripten_resize_heap:_emscripten_resize_heap,random_get:_random_get};function run(){preRun();function doRun(){Module["calledRun"]=true;if(ABORT)return;initRuntime();readyPromiseResolve?.(Module);Module["onRuntimeInitialized"]?.();postRun()}if(Module["setStatus"]){Module["setStatus"]("Running...");setTimeout(()=>{setTimeout(()=>Module["setStatus"](""),1);doRun()},1)}else{doRun()}}var wasmExports;wasmExports=await (createWasm());run();if(runtimeInitialized){moduleRtn=Module}else{moduleRtn=new Promise((resolve,reject)=>{readyPromiseResolve=resolve;readyPromiseReject=reject})} ;return moduleRtn}export default NispsModule; diff --git a/playground/wasm/nisps.wasm b/playground/wasm/nisps.wasm index 28e296c418628fa0e05d02834405961848001ab1..324156207dd0332c4c1604a382227cb4f489b069 100755 GIT binary patch delta 6850 zcmZ`;du$xXd7qhm+}`c(-QI)b@$N{?E-9JR!%||IqGXA@w!a%#zuZCDheFtP4P zN~9=DHbSq0xOEYyK+UC)0&UW?hS9=^8%I~5D4I42@=uF6h>NC4)1(g21Wk&xK!XMa zYS>VJ-|Uf;Wt76r%+BL`zh?AbFNz<1QCwiNR~I;AjK9jeF7Q>pdVxngLc=a-QG`}U zi;l4i28}Q@Ro`Bs6M|;Z;2fPn$MB)618yGwtQCJaQ_dJ$Ba64MSwOFw?G<&Ey16r%f+kNEJ&R8W(BC zG%&FFXK5D3oesF-f*ah%j||7GczBXUAg!)s}=4!H8XdE^>lAdE?jD%p@4^O!<&eIrO)BW zM~xk5pEp7r{D#rP3-Q~=af}DeRkVL%{w7xUr25c)M77saduyY%WCOwpxL%h`@*&~T zcc~(|HyVEh!VBD7vdccvh}H1& z1Xrs|osJnmp-Y3d!gO?VZsP8byriWn?y5R)v0veC~-tgWaJa0vR3i$8)|o-e(~y_VT=j;3Zh zXV89KwcEOGV*I79pP;>2eiU23SKf~CJLRKj$Gb=QL-A+OPR8Hr{vpO+sJu$^)qc!8 zU0p%@FV$^0^p`(N9Vb~BPEQfphI4i}mY9OT1-YC4j*6+&>gHN=+x_w2! z+jb1>A%F5nG&)4%Nsv&cZ|F&fKV2TEfXvShtfB20B(DlpJ2&WK{Nms#wBJ`V@2IwS z=OW(`|IW_sXj`iNp=$4_cKa@xKS6D4Y1d1fA8Ng~dw_(y$?Ur}T*3Ax{GwEo#m&jAEHr#i zW@WMAX8~KR1Zb}L-BP15nrnUqhF^%69t?OfUVrd{U6=~01lq&wH-}1vi6BFk>B@|> z8{`0m_~%3WA%oh1k@)w9YA&Ep!H97Lqpc(0y?sPtdh-a+#!nuoM45gjUE-vRqCmGq z$ASuYTkvf-uobe0T{h0x%~5H$Zu#3nF2Q5SPgeb!Ru}Dmbc|w%I;~CV^^o~CekBK% zW=gFYKyVA;_h0+R*Eo!dgcOu~TdSu53$enIcKqIfzK9eu5#-67ZHW#kO){Bp1P<+C zU@Z?1Wl7MnC1!P^#%cIABnd7wX->QFXyDF+e9+?s9n?7iOUTl&gp$FQ9g;O6`Mb%A zG+^)aHE06c2&F~=O~o!f0Ado3Ng+~~W;4(!MCug0aExfd5FMi#31V^tf!xD^4D^I( zO=1>A6Scl384=B7ha!5s$nfX)%5GVias{Z^=J&)OI#`T)Y?&c~N-{$PIkF`4p)()Q zVUqj@NMl5r2ECI!;hP?Y=p34nx+1Nl6;L3}W{{K4fEa-7pz9Ahn$QyV zDqfw0R}DW!09Nl<{Ci<$t8y!X2@htIiVvl305MIlpqHgJ*2<>yv2Ns z6g_A#pYjw9j*4Um{Uj|AgUSGr3`RB}%7IDXp(uRR0v+T)IXcAxWfFuO39Ly1gm^=} z3DzngbUV&NzBmD+RFHh?CCIPQr>sx}iobcda*;5F`;bm#k0b2x^KKZY!~fL7AKM-> zEG2G`O(F}(4WJCNH+)+P1hl4Xd}9gvRvZAjEqi6p8my`oj~wZ;d*wFSyDrOBnAU|O zrM?WM*ct%#3{&ESoz)s>$_-`h;61O!pFi?I1UQTgR(2{|m7N5+1bk={dB6|hP#YB5 zBo7b?7v7MEBBms@z)nDv2t-=;oJT}JDc1-LX(>TjzF||5P!DXy5GWknB#{hsgG3wH z6p4^(Mn!Dd6R(VnB98z02xWtR8W};`HA>@$Mk~`mSOiFllF(q5C{{8hJXxaHCW^Bg z6x&2$1|d5miwVs|nkmG8FuDf?6_&`c3dvFdn*8zTr;wZ;ucvwXke~Jr`5Es}+zh`G z{g1|-sE>uR8JN&9OqHZ=l7*K9!VU0i;Z6MM6dw^XDKA*C^fXzwDJ@iwN$EvNEe%Vz zl%<>P%1Z(C5`Y!}3Kh5$Fm@=SYrY4DO)7pUf;)kR;)_Na37#r3&ohAsl_T;c{5aWb zqM8U%*(<*_Xcfjw#gf%T&I;1fkm;k+Z4!G-N^2z;aF5Cmg+#mTCe1-zRAXQXHHM~^ zHY-DBziA5Gi?vl~e;*QIKO^-JM3E%KV}Kwin!6|+a+HmmAOW_ySz_cVsU!ZiOAOlQ zNj2$hPG)7UfdVBfy@sDpN{lR(74SD`B1-Cw4zQ14N60i%zadHUswc)>e00DJGytI# zCbO8(Q9n>$!h#z(QjkcHKq(}Fle;SffeN9l)ln`m=68|m+pKC)U6%b$slq5&gK0p~ zekUbwIKpmelOyyeUo!ckV3Z>oY-AC$r-l)^nw3V@PJpyU+KNs27m?+D>W`YLHhN8V6rSXE|X{5x;H`Pb{JWL@)% zN$D!07>3wizP3Av18PS1y zeSKYatji2Au942$NtY-i(THou;Mz#ZWC|34hmmqNVxdI=0(H7lhW&~z%*~LN>Tll= z#O+&t4iXVbIv_3qN5vl>67vl(pF@I3!T>F+FaYaD0--RFQeglPD1iVQIznsDlSTl6 z+&P5^tL{SmG$awh`F5)eO0vT3S4&Fm6nW)90y`2i0bZtzH4lry%|k&&hDxBO0aaZU zkw+-m;KCX7`#!Yh+<)x=Rf<%YOcNpZXypb$7tK|z_cpGrD2>2z2PMBoej(t^TJM!Bo< zcT}Y0t_55flHFR;8HjbTG{i&P`ly5dQS=2qioKK-sU~57=Y~|wBTzoSswRDV7J=yQOl=`kr;t)0P93b0-DwCy%8;dj5pHx+ABGvB5R@s3PuL1q z%uqFjAYrGWK8VR1u}zFn`tBU_@fh*&SP$Mo8)Nrz_25cwt<%FJ9vyqzQG$KW8;)Ne z@5PJeo8#ZZnLnCX$Me(BvelozZfp;qKP0vr^FYrc!hpI16{xW+6bh7n0L3pI~ z>J@2#*}6_pU!X-evUJ;R}`cjmZaE?;ZbF9t^KZ9m?K3Q@Qop^ zkr>AOl`0{dukUy5dZqrz~sU3PHrh6!M`!)9Yl~+b|w`+cb;^4 z1>ZL#@1WF9Dl21aIoRfS*Ya1#&L_5c#T$xePZ#-6>)L75!iNb!eeG84b9#yp;s1X2 zrC)1Q9&^dpyu)$n(LR2-HT-D9adkyj@$xB|Z7|f3whqbK zSyJg+tl^Ec9$WdG;ClSiXD0FHICy=HpNPM9y;?aze33uloe?&V>mcx;>dou%udct& zkHvp@V-W4%-S|l|){mjijlX@PC;sne&$(3SCdYK|RDABaUH8oV?sJDI$HThFqRyy` LZR?+(>rDS2sw9?+ delta 3788 zcmaJ@eQX>@6`$GNTkq}e-R*tZ=ex6ScI~u|oew)s&L2%|Z&Eu>Q^!uylqe-RN9ncA zUF@W`Lla22ex(H!wIq|O5~2vGK&Yi^#8LH+B1EY`s)Ps1ZVlp+xKSP`@J`B-pspSQ~vz8a*!Qeoa2l!el9u7&-3%MyvA!FY*wq`$7a=l zGBz8cZwT~Y1f$Rw2w6t{J!pR4icJZB= z$*?*z9a2LfCCnAw&?4ce!VN=@IdQHSjA;OL<7a9n_=F9VT;U;Z;h!s(rKzeaAK*8s zFRN+!b-p{Dh|eEiJhfPwJ8`OXcX@GX;pADCyrg7WQ|VR3z2&9S1LfoQ-nX>K(y~u! zS=*dk)gGU}r@T;}KUQWfu~q5P!qMaNEF5{gET_*km_f54 zzGtCaW_c|bhL#NlGfKzjmsl>kIx7(5FBX)sbO}X zyrQ-j_lJ8p^SQNVEWCcB-Cl6Lj?OhV-Q2adyJy{=vHbS!x88o+M1J?Ko%yNhnQ5CV z=662i(^AuZ$cNxx<+?$F-VoIKF%+lr*+@+>cb8 z5&6|n87me2CeXFIkM&i(jmPB|_1nO|shX8^O;QSztuSYkJRHkFomo-a3tkhco0ip2#U}YC`F8AC82^m(8`wJ# zKLh^faS!yPcn&m@cm?vGBG#s_V#brrkMmf)x#b{-JcufTciXXmA}lr1OCPQw@BXB3CWqZCD0Gr z3Iw5j7ay&Es{Ik3-5|KH6%`Ewrw!x2K>=}RgZyjZFyAWoiBe{(DdGNejqW|k#@tnc`oF$ z>(+q29?(AoG`5~3TLO9)(fUK{pWu9RJ+iTjl3VP38moWSI{}*PBTEB)TM1f(pc#M` z5e7lSH=r4S_Eg_{7}sT;zWNLE}+s% z#$a7p@<2|ztwJRphgPR)`YO$RWgaTbr z$sRz3YcLe>#*+SQRGOe-pe(4{rpk<_t= zBVL?NfG#jF$o=8$0bSUTS(t2BTpd{kU{#nDmIXI)AO!-)OVD6@tbssHz+zw|5%3o!m~=jDxl9NQ%!?9;gosu>_?g5BL;(Ov;2o?;RJ?YwVK$}- zvyntiK$-y*GL2gBSq;$O0BBMTQ@E#ScTIomiHgWKSR1V7rz)ZegHdrDvWt2^(o!A; ztI_6U=o?)v>Y!l-8AYp032m)nyI$Z@4P4TGYr-4A0@hUvd~u57|2w3ELJkbi9FHJJ z)Cz18I9Y^F9)l19Fi@?4X{dtkiE~7^BkipnToYJc%m54oso6b}v@co}30aQNlzX=&YjBk(10qY!J223B!?@B#s&y;2{a*BNQLq)ueXPP z|JrhJ3<(zs!J@6DYDGQM@x`~XY!@uL!iow9Ccw;W@+m~&k}E`;{BT>TcU!w>WDW;g6?c1>rq z%x-_l3FG9TB|3wa4h}P)o^o+}y)xB?*Wg=IH`8tY&plbZ=Q7jnpaauyLT%4n##`qb zGxRjNKGVY6WXs+^MF%{XxGWo%2I|hfG@b;Qo0r=UoWaaj4;*gA+bD!ghwz9(f2X&R z|6o?AcpUZJ*{XP(E(xt0HDpVt!5XWq9kBAm!Cv5a_TU}7um17DF^-S^Lp{7tmJW4w z_50Uj_j@YtfGgXb{$0Tp%7BV)((89cc$S!&(=T5-lu!4@P(3@o|IyVSAI9ONd&6;h z>%TozQ+ZL&9%%=y9qHi3`g2DrIv9g808+6=$&iTJ1{;SLL3E$MI~v{O~1h+_E;y=p&izIDx0eg4e20kWy-Z=wxwK ze)^FY`3`A5x>ntxjY #include +// Helper subclass to access protected members for extended training +struct NispsMLPAccessor : public nisps::MLP { + using nisps::MLP::loss_fn_; + using nisps::MLP::UpdateWeights; +}; + extern "C" { // ---- Lifecycle ---- @@ -163,6 +169,182 @@ void nisps_mlp_move_weights_spread(void* ptr, float speed, float spread) { } } +// ---- Batch inference ---- + +EMSCRIPTEN_KEEPALIVE +void nisps_mlp_infer_batch(void* ptr, float* inputs_flat, int n_points, int input_dim, float* outputs_flat, int output_dim) { + auto* mlp = static_cast*>(ptr); + std::vector in_vec(input_dim); + std::vector out_vec; + for (int i = 0; i < n_points; i++) { + in_vec.assign(inputs_flat + i * input_dim, inputs_flat + (i + 1) * input_dim); + out_vec.clear(); + mlp->GetOutput(in_vec, &out_vec, nullptr, true); + int n = output_dim < (int)out_vec.size() ? output_dim : (int)out_vec.size(); + for (int j = 0; j < n; j++) { + outputs_flat[i * output_dim + j] = out_vec[j]; + } + } +} + +// ---- Extended training with per-iteration loss history ---- + +EMSCRIPTEN_KEEPALIVE +int nisps_mlp_train_ex(void* ptr, + float* features_flat, int n_samples, int feature_dim, + float* labels_flat, int label_dim, + float* sample_weights, + float learning_rate, int max_iterations, float min_error, + float* loss_history_out) { + + auto* mlp = static_cast(static_cast*>(ptr)); + + std::vector> features(n_samples); + std::vector> labels(n_samples); + + for (int i = 0; i < n_samples; i++) { + features[i].assign(features_flat + i * feature_dim, + features_flat + (i + 1) * feature_dim); + labels[i].assign(labels_flat + i * label_dim, + labels_flat + (i + 1) * label_dim); + } + + float sample_size_recip = 1.0f / n_samples; + + int iter = 0; + for (iter = 0; iter < max_iterations; iter++) { + float iteration_loss = 0.0f; + + for (int s = 0; s < n_samples; s++) { + float w = sample_weights ? sample_weights[s] : sample_size_recip; + + std::vector predicted_output; + std::vector> all_layers_activations; + + mlp->GetOutput(features[s], &predicted_output, &all_layers_activations, false); + + std::vector deriv_error_output(predicted_output.size()); + float loss = mlp->loss_fn_(labels[s], predicted_output, deriv_error_output, w); + + iteration_loss += loss; + + mlp->UpdateWeights(all_layers_activations, deriv_error_output, learning_rate); + } + + if (!sample_weights) { + iteration_loss *= sample_size_recip; + } + + loss_history_out[iter] = iteration_loss; + + if (iteration_loss < min_error) { + iter++; + break; + } + } + + return iter; +} + +// ---- moveWeights with output pin mask ---- + +EMSCRIPTEN_KEEPALIVE +void nisps_mlp_move_weights_ex(void* ptr, float speed, float spread, int* pin_mask, int n_outputs) { + auto* mlp = static_cast*>(ptr); + float decay = 1.0f - 0.1f * spread; + size_t n_layers = mlp->m_layers.size(); + for (size_t l = 0; l < n_layers; l++) { + int fan_in = mlp->m_layers[l].GetInputSize(); + float xavier_scale = 1.0f / std::sqrt((float)fan_in); + float layer_scale = 1.0f * (1.0f - spread) + xavier_scale * spread; + bool is_output_layer = (l == n_layers - 1); + int node_idx = 0; + for (auto& node : mlp->m_layers[l].GetNodesChangeable()) { + // Skip pinned output nodes + if (is_output_layer && pin_mask && node_idx < n_outputs && pin_mask[node_idx] == 1) { + node_idx++; + continue; + } + for (size_t j = 0; j < node.m_weights.size(); j++) { + node.m_weights[j] *= decay; + float accum = 0; + for (int n = 0; n < 3; n++) { + accum += (float)rand() / RAND_MAX * 2.0f - 1.0f; + } + node.m_weights[j] += 3.0f * accum * speed * layer_scale; + } + node_idx++; + } + } +} + +// ---- Evaluate loss without updating weights ---- + +EMSCRIPTEN_KEEPALIVE +float nisps_mlp_eval_loss(void* ptr, + float* features_flat, int n_samples, int feature_dim, + float* labels_flat, int label_dim, + float* sample_weights) { + + auto* mlp = static_cast(static_cast*>(ptr)); + float total_loss = 0.0f; + float sample_size_recip = 1.0f / n_samples; + + for (int s = 0; s < n_samples; s++) { + float w = sample_weights ? sample_weights[s] : sample_size_recip; + + std::vector in_vec(features_flat + s * feature_dim, + features_flat + (s + 1) * feature_dim); + std::vector label_vec(labels_flat + s * label_dim, + labels_flat + (s + 1) * label_dim); + + std::vector predicted_output; + mlp->GetOutput(in_vec, &predicted_output, nullptr, true); + + std::vector deriv_error_output(predicted_output.size()); + float loss = mlp->loss_fn_(label_vec, predicted_output, deriv_error_output, w); + total_loss += loss; + } + + if (!sample_weights) { + total_loss *= sample_size_recip; + } + + return total_loss; +} + +// ---- Per-layer weight statistics ---- + +EMSCRIPTEN_KEEPALIVE +void nisps_mlp_get_layer_stats(void* ptr, float* stats_out, int n_layers) { + auto* mlp = static_cast*>(ptr); + int layers_to_process = n_layers < (int)mlp->m_layers.size() ? n_layers : (int)mlp->m_layers.size(); + for (int l = 0; l < layers_to_process; l++) { + float sum_abs = 0.0f; + float max_abs = 0.0f; + int dead_count = 0; + int saturating_count = 0; + int total_weights = 0; + + for (auto& node : mlp->m_layers[l].m_nodes) { + for (size_t j = 0; j < node.m_weights.size(); j++) { + float aw = std::fabs(node.m_weights[j]); + sum_abs += aw; + if (aw > max_abs) max_abs = aw; + if (aw < 0.01f) dead_count++; + if (aw > 3.0f) saturating_count++; + total_weights++; + } + } + + float inv_total = total_weights > 0 ? 1.0f / total_weights : 0.0f; + stats_out[l * 4 + 0] = sum_abs * inv_total; // mean absolute weight + stats_out[l * 4 + 1] = max_abs; // max absolute weight + stats_out[l * 4 + 2] = dead_count * inv_total; // fraction dead + stats_out[l * 4 + 3] = saturating_count * inv_total; // fraction saturating + } +} + // ---- Memory helpers ---- EMSCRIPTEN_KEEPALIVE