feat(playground/wasm): wire 5 new WASM functions into JS WasmIML + worker

Expose inferBatch, trainEx, moveWeightsEx, evalLoss, and getLayerStats
from the WASM binary into the JavaScript layer:

- inferBatch: batch inference for heatmap sampling
- trainEx: replaces train() with per-iteration loss history capture
- moveWeightsEx: native output pin mask support (removes save/restore hack)
- evalLoss: compute loss without updating weights
- getLayerStats: per-layer weight statistics (meanAbs, maxAbs, dead/sat fracs)

Worker also upgraded to trainEx, returning full lossHistory in payload.
This commit is contained in:
w1n5t0n 2026-04-03 17:15:38 +01:00
parent f225c7c2a6
commit 7f90b4323d
2 changed files with 145 additions and 45 deletions

View file

@ -19,6 +19,7 @@ async function ensureModule() {
getWeights: mod.cwrap('nisps_mlp_get_weights', null, ['number', 'number']), getWeights: mod.cwrap('nisps_mlp_get_weights', null, ['number', 'number']),
setWeights: mod.cwrap('nisps_mlp_set_weights', null, ['number', 'number']), setWeights: mod.cwrap('nisps_mlp_set_weights', null, ['number', 'number']),
train: mod.cwrap('nisps_mlp_train', 'number', ['number', 'number', 'number', 'number', 'number', 'number', 'number', 'number', 'number', 'number']), train: mod.cwrap('nisps_mlp_train', 'number', ['number', 'number', 'number', 'number', 'number', 'number', 'number', 'number', 'number', 'number']),
trainEx: mod.cwrap('nisps_mlp_train_ex', 'number', ['number', 'number', 'number', 'number', 'number', 'number', 'number', 'number', 'number', 'number', 'number']),
alloc: mod.cwrap('nisps_alloc', 'number', ['number']), alloc: mod.cwrap('nisps_alloc', 'number', ['number']),
free: mod.cwrap('nisps_free', null, ['number']), free: mod.cwrap('nisps_free', null, ['number']),
allocInt: mod.cwrap('nisps_alloc_int', 'number', ['number']), allocInt: mod.cwrap('nisps_alloc_int', 'number', ['number']),
@ -98,18 +99,26 @@ self.onmessage = async function(e) {
const featPtr = toHeapF32(featFlat); const featPtr = toHeapF32(featFlat);
const labPtr = toHeapF32(labFlat); const labPtr = toHeapF32(labFlat);
const weightPtr = sampleWeights ? toHeapF32(new Float32Array(sampleWeights)) : 0; const weightPtr = sampleWeights ? toHeapF32(new Float32Array(sampleWeights)) : 0;
const lossHistPtr = w.alloc(maxIterations);
// Train // Train (extended — captures per-iteration loss history)
const loss = w.train( const itersRun = w.trainEx(
mlp, featPtr, nSamples, featureDim, mlp, featPtr, nSamples, featureDim,
labPtr, nOutputs, labPtr, nOutputs,
weightPtr, weightPtr,
learningRate, maxIterations, convergenceThreshold learningRate, maxIterations, convergenceThreshold,
lossHistPtr
); );
// Read per-iteration loss history
const lossHistory = fromHeapF32(lossHistPtr, itersRun);
w.free(featPtr); w.free(featPtr);
w.free(labPtr); w.free(labPtr);
if (weightPtr) w.free(weightPtr); if (weightPtr) w.free(weightPtr);
w.free(lossHistPtr);
const loss = itersRun > 0 ? lossHistory[itersRun - 1] : 0;
// Extract trained weights // Extract trained weights
const outPtr = w.alloc(weightCount); const outPtr = w.alloc(weightCount);
@ -119,7 +128,7 @@ self.onmessage = async function(e) {
self.postMessage({ self.postMessage({
type: 'trained', type: 'trained',
payload: { weights: trainedWeights, loss }, payload: { weights: trainedWeights, loss, lossHistory },
}); });
} }
}; };

View file

@ -31,6 +31,11 @@ function wrapModule(mod) {
train: mod.cwrap('nisps_mlp_train', 'number', ['number', 'number', 'number', 'number', 'number', 'number', 'number', 'number', 'number', 'number']), train: mod.cwrap('nisps_mlp_train', 'number', ['number', 'number', 'number', 'number', 'number', 'number', 'number', 'number', 'number', 'number']),
drawWeightsSpread: mod.cwrap('nisps_mlp_draw_weights_spread', null, ['number', 'number']), drawWeightsSpread: mod.cwrap('nisps_mlp_draw_weights_spread', null, ['number', 'number']),
moveWeightsSpread: mod.cwrap('nisps_mlp_move_weights_spread', null, ['number', 'number', 'number']), moveWeightsSpread: mod.cwrap('nisps_mlp_move_weights_spread', null, ['number', 'number', 'number']),
inferBatch: mod.cwrap('nisps_mlp_infer_batch', null, ['number', 'number', 'number', 'number', 'number', 'number']),
trainEx: mod.cwrap('nisps_mlp_train_ex', 'number', ['number', 'number', 'number', 'number', 'number', 'number', 'number', 'number', 'number', 'number', 'number']),
moveWeightsEx: mod.cwrap('nisps_mlp_move_weights_ex', null, ['number', 'number', 'number', 'number', 'number']),
evalLoss: mod.cwrap('nisps_mlp_eval_loss', 'number', ['number', 'number', 'number', 'number', 'number', 'number', 'number']),
getLayerStats: mod.cwrap('nisps_mlp_get_layer_stats', null, ['number', 'number', 'number']),
alloc: mod.cwrap('nisps_alloc', 'number', ['number']), alloc: mod.cwrap('nisps_alloc', 'number', ['number']),
free: mod.cwrap('nisps_free', null, ['number']), free: mod.cwrap('nisps_free', null, ['number']),
allocInt: mod.cwrap('nisps_alloc_int', 'number', ['number']), allocInt: mod.cwrap('nisps_alloc_int', 'number', ['number']),
@ -196,6 +201,38 @@ export class WasmIML {
this.inputUpdated = false; this.inputUpdated = false;
} }
// ---- Batch inference (WASM) ----
inferBatch(inputPoints) {
// inputPoints: array of [x,y,...] arrays (each length nInputs)
// Returns: array of output arrays (each length nOutputs)
const nPoints = inputPoints.length;
const inputDim = this.nInputs + 1; // +bias
const inFlat = new Float32Array(nPoints * inputDim);
for (let i = 0; i < nPoints; i++) {
for (let j = 0; j < this.nInputs; j++) {
inFlat[i * inputDim + j] = inputPoints[i][j];
}
inFlat[i * inputDim + this.nInputs] = 1.0; // bias
}
const inPtr = toHeapF32(this._w, inFlat);
const outPtr = this._w.alloc(nPoints * this.nOutputs);
this._w.inferBatch(this._mlp, inPtr, nPoints, inputDim, outPtr, this.nOutputs);
// Read outputs
const results = [];
const heap = this._w.mod.HEAPF32;
const outOff = outPtr >> 2;
for (let i = 0; i < nPoints; i++) {
const row = new Array(this.nOutputs);
for (let j = 0; j < this.nOutputs; j++) {
row[j] = heap[outOff + i * this.nOutputs + j];
}
results.push(row);
}
this._w.free(inPtr);
this._w.free(outPtr);
return results;
}
// ---- Dataset ---- // ---- Dataset ----
addExample(inputs, outputs) { addExample(inputs, outputs) {
const inVec = inputs.slice(0, this.nInputs); const inVec = inputs.slice(0, this.nInputs);
@ -255,23 +292,30 @@ export class WasmIML {
const featPtr = toHeapF32(this._w, featFlat); const featPtr = toHeapF32(this._w, featFlat);
const labPtr = toHeapF32(this._w, labFlat); const labPtr = toHeapF32(this._w, labFlat);
const weightPtr = toHeapF32(this._w, sampleWeights); const weightPtr = toHeapF32(this._w, sampleWeights);
const lossHistPtr = this._w.alloc(this.maxIterations);
const loss = this._w.train( const itersRun = this._w.trainEx(
this._mlp, featPtr, nSamples, featureDim, this._mlp, featPtr, nSamples, featureDim,
labPtr, this.nOutputs, labPtr, this.nOutputs,
weightPtr, weightPtr,
this.learningRate, this.maxIterations, this.convergenceThreshold this.learningRate, this.maxIterations, this.convergenceThreshold,
lossHistPtr
); );
// Read per-iteration loss history
const lossHist = fromHeapF32(this._w, lossHistPtr, itersRun);
this._w.free(featPtr); this._w.free(featPtr);
this._w.free(labPtr); this._w.free(labPtr);
this._w.free(weightPtr); this._w.free(weightPtr);
this._w.free(lossHistPtr);
const loss = itersRun > 0 ? lossHist[itersRun - 1] : 0;
this.lastLoss = loss; this.lastLoss = loss;
// We don't have per-iteration history from WASM (single return value), for (let i = 0; i < lossHist.length; i++) {
// so record just the final loss this.lossHistory.push(lossHist[i]);
this.lossHistory.push(loss); }
this.totalTrainingIterations += 1; this.totalTrainingIterations += itersRun;
if (this.lossHistory.length > 1200) { if (this.lossHistory.length > 1200) {
this.lossHistory = this.lossHistory.slice(this.lossHistory.length - 1200); this.lossHistory = this.lossHistory.slice(this.lossHistory.length - 1200);
} }
@ -285,6 +329,51 @@ export class WasmIML {
return loss; return loss;
} }
// ---- Loss evaluation (WASM, no weight update) ----
evalLoss() {
if (this.dataset.features.length === 0) return null;
const features = this.dataset.features;
const labels = this.dataset.labels;
const sampleWeights = this.dataset.computeWeights(this.weightingMode, {
recencyBias: this.recencyBias,
queryInput: this.inputState,
radius: this.localRadius,
});
const featureDim = this.nInputs + 1;
const nSamples = features.length;
const featFlat = new Float32Array(nSamples * featureDim);
const labFlat = new Float32Array(nSamples * this.nOutputs);
for (let i = 0; i < nSamples; i++) {
for (let j = 0; j < this.nInputs; j++) {
featFlat[i * featureDim + j] = features[i][j];
}
featFlat[i * featureDim + this.nInputs] = 1.0;
for (let j = 0; j < this.nOutputs; j++) {
labFlat[i * this.nOutputs + j] = labels[i][j] || 0;
}
}
const featPtr = toHeapF32(this._w, featFlat);
const labPtr = toHeapF32(this._w, labFlat);
const weightPtr = toHeapF32(this._w, sampleWeights);
const loss = this._w.evalLoss(
this._mlp, featPtr, nSamples, featureDim,
labPtr, this.nOutputs,
weightPtr
);
this._w.free(featPtr);
this._w.free(labPtr);
this._w.free(weightPtr);
return loss;
}
// ---- Weight manipulation ---- // ---- Weight manipulation ----
randomiseWeights(spread = 0) { randomiseWeights(spread = 0) {
this.storedWeights = this._getFlatWeights(); this.storedWeights = this._getFlatWeights();
@ -297,45 +386,20 @@ export class WasmIML {
} }
// outputPinMask: optional Uint8Array[nOutputs], 1 = skip that output node. // outputPinMask: optional Uint8Array[nOutputs], 1 = skip that output node.
// Since WASM moveWeights doesn't support pin masks, we save pinned nodes'
// weights before the call and restore them after.
moveWeights(speed, spread = 0, outputPinMask = null) { moveWeights(speed, spread = 0, outputPinMask = null) {
let savedSlices = null; let pinMaskPtr = 0;
if (outputPinMask && outputPinMask.some(v => v)) { if (outputPinMask && outputPinMask.some(v => v)) {
// Compute the flat-array offset of the last layer's nodes. const pinI32 = new Int32Array(this.nOutputs);
// Flat format: for each layer, for each node: [w0..wN, bias] for (let i = 0; i < this.nOutputs; i++) {
const allWeights = this._getFlatWeights(); pinI32[i] = outputPinMask[i] ? 1 : 0;
const lastLayerInputSize = this.layerSizes[this.layerSizes.length - 2];
const numOutputNodes = this.layerSizes[this.layerSizes.length - 1];
const weightsPerOutputNode = lastLayerInputSize + 1; // weights + bias
// Offset of the last layer in the flat array
const lastLayerOffset = this._weightCount - (numOutputNodes * weightsPerOutputNode);
// Save pinned nodes' weight slices
savedSlices = [];
for (let i = 0; i < numOutputNodes; i++) {
if (outputPinMask[i]) {
const start = lastLayerOffset + i * weightsPerOutputNode;
const end = start + weightsPerOutputNode;
savedSlices.push({ start, end, data: allWeights.slice(start, end) });
}
} }
pinMaskPtr = toHeapI32(this._w, pinI32);
} }
this._w.moveWeightsSpread(this._mlp, speed, spread); this._w.moveWeightsEx(this._mlp, speed, spread, pinMaskPtr, this.nOutputs);
// Restore pinned nodes' weights if (pinMaskPtr) this._w.freeInt(pinMaskPtr);
if (savedSlices && savedSlices.length > 0) {
const allWeights = this._getFlatWeights();
for (const slice of savedSlices) {
for (let j = 0; j < slice.data.length; j++) {
allWeights[slice.start + j] = slice.data[j];
}
}
this._setFlatWeights(allWeights);
}
this.inputUpdated = true; this.inputUpdated = true;
this.process(); this.process();
@ -356,6 +420,26 @@ export class WasmIML {
this._w.free(ptr); this._w.free(ptr);
} }
// ---- Per-layer weight statistics (WASM) ----
getLayerStats() {
const nLayers = this.layerSizes.length - 1;
const statsPtr = this._w.alloc(nLayers * 4);
this._w.getLayerStats(this._mlp, statsPtr, nLayers);
const stats = [];
const heap = this._w.mod.HEAPF32;
const off = statsPtr >> 2;
for (let l = 0; l < nLayers; l++) {
stats.push({
meanAbs: heap[off + l * 4 + 0],
maxAbs: heap[off + l * 4 + 1],
deadFrac: heap[off + l * 4 + 2],
satFrac: heap[off + l * 4 + 3],
});
}
this._w.free(statsPtr);
return stats;
}
// ---- Async training via Web Worker ---- // ---- Async training via Web Worker ----
get isTraining() { return this._training; } get isTraining() { return this._training; }
@ -401,13 +485,20 @@ export class WasmIML {
this._worker.removeEventListener('message', handler); this._worker.removeEventListener('message', handler);
this._training = false; this._training = false;
const { weights, loss } = e.data.payload; const { weights, loss, lossHistory } = e.data.payload;
// Swap in trained weights // Swap in trained weights
this._setFlatWeights(weights); this._setFlatWeights(weights);
this.lastLoss = loss; this.lastLoss = loss;
if (lossHistory && lossHistory.length > 0) {
for (let i = 0; i < lossHistory.length; i++) {
this.lossHistory.push(lossHistory[i]);
}
this.totalTrainingIterations += lossHistory.length;
} else {
this.lossHistory.push(loss); this.lossHistory.push(loss);
this.totalTrainingIterations += 1; this.totalTrainingIterations += 1;
}
if (this.lossHistory.length > 1200) { if (this.lossHistory.length > 1200) {
this.lossHistory = this.lossHistory.slice(this.lossHistory.length - 1200); this.lossHistory = this.lossHistory.slice(this.lossHistory.length - 1200);
} }