2026-03-23 22:23:07 +01:00
|
|
|
// Web Worker for async WASM training (module worker)
|
|
|
|
|
// Loads its own nisps WASM instance for off-thread training
|
|
|
|
|
|
|
|
|
|
import NispsModule from '../../wasm/nisps.js';
|
|
|
|
|
|
|
|
|
|
let mod = null;
|
|
|
|
|
let w = null;
|
|
|
|
|
let mlp = null;
|
|
|
|
|
let currentLayerSizes = null;
|
|
|
|
|
|
|
|
|
|
async function ensureModule() {
|
|
|
|
|
if (mod) return;
|
|
|
|
|
mod = await NispsModule();
|
|
|
|
|
w = {
|
|
|
|
|
mod,
|
|
|
|
|
create: mod.cwrap('nisps_mlp_create', 'number', ['number', 'number', 'number', 'number']),
|
|
|
|
|
destroy: mod.cwrap('nisps_mlp_destroy', null, ['number']),
|
|
|
|
|
weightCount: mod.cwrap('nisps_mlp_weight_count', 'number', ['number']),
|
|
|
|
|
getWeights: mod.cwrap('nisps_mlp_get_weights', null, ['number', 'number']),
|
|
|
|
|
setWeights: mod.cwrap('nisps_mlp_set_weights', null, ['number', 'number']),
|
2026-04-02 21:35:31 +02:00
|
|
|
train: mod.cwrap('nisps_mlp_train', 'number', ['number', 'number', 'number', 'number', 'number', 'number', 'number', 'number', 'number', 'number']),
|
2026-04-03 18:15:38 +02:00
|
|
|
trainEx: mod.cwrap('nisps_mlp_train_ex', 'number', ['number', 'number', 'number', 'number', 'number', 'number', 'number', 'number', 'number', 'number', 'number']),
|
2026-03-23 22:23:07 +01:00
|
|
|
alloc: mod.cwrap('nisps_alloc', 'number', ['number']),
|
|
|
|
|
free: mod.cwrap('nisps_free', null, ['number']),
|
|
|
|
|
allocInt: mod.cwrap('nisps_alloc_int', 'number', ['number']),
|
|
|
|
|
freeInt: mod.cwrap('nisps_free_int', null, ['number']),
|
|
|
|
|
};
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
function toHeapF32(arr) {
|
|
|
|
|
const ptr = w.alloc(arr.length);
|
|
|
|
|
w.mod.HEAPF32.set(arr, ptr >> 2);
|
|
|
|
|
return ptr;
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
function toHeapI32(arr) {
|
|
|
|
|
const ptr = w.allocInt(arr.length);
|
|
|
|
|
w.mod.HEAP32.set(arr, ptr >> 2);
|
|
|
|
|
return ptr;
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
function fromHeapF32(ptr, length) {
|
|
|
|
|
const offset = ptr >> 2;
|
|
|
|
|
return Array.from(w.mod.HEAPF32.subarray(offset, offset + length));
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
function arraysEqual(a, b) {
|
|
|
|
|
if (!a || !b || a.length !== b.length) return false;
|
|
|
|
|
for (let i = 0; i < a.length; i++) if (a[i] !== b[i]) return false;
|
|
|
|
|
return true;
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
self.onmessage = async function(e) {
|
|
|
|
|
const { type, payload } = e.data;
|
|
|
|
|
|
|
|
|
|
if (type === 'train') {
|
|
|
|
|
await ensureModule();
|
|
|
|
|
|
|
|
|
|
const {
|
|
|
|
|
layerSizes, activationIds, weights,
|
2026-04-02 21:35:31 +02:00
|
|
|
features, labels, sampleWeights,
|
2026-03-23 22:23:07 +01:00
|
|
|
nInputs, nOutputs,
|
|
|
|
|
learningRate, maxIterations, convergenceThreshold,
|
|
|
|
|
} = payload;
|
|
|
|
|
|
|
|
|
|
// Recreate MLP if architecture changed
|
|
|
|
|
if (!mlp || !arraysEqual(currentLayerSizes, layerSizes)) {
|
|
|
|
|
if (mlp) w.destroy(mlp);
|
|
|
|
|
const layerPtr = toHeapI32(new Int32Array(layerSizes));
|
|
|
|
|
const actPtr = toHeapI32(new Int32Array(activationIds));
|
|
|
|
|
mlp = w.create(layerPtr, layerSizes.length, actPtr, activationIds.length);
|
|
|
|
|
w.freeInt(layerPtr);
|
|
|
|
|
w.freeInt(actPtr);
|
|
|
|
|
currentLayerSizes = [...layerSizes];
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
// Load weights from main thread
|
|
|
|
|
const weightCount = w.weightCount(mlp);
|
|
|
|
|
const wPtr = toHeapF32(new Float32Array(weights));
|
|
|
|
|
w.setWeights(mlp, wPtr);
|
|
|
|
|
w.free(wPtr);
|
|
|
|
|
|
|
|
|
|
// Build flat training arrays with bias
|
|
|
|
|
const featureDim = nInputs + 1;
|
|
|
|
|
const nSamples = features.length;
|
|
|
|
|
const featFlat = new Float32Array(nSamples * featureDim);
|
|
|
|
|
const labFlat = new Float32Array(nSamples * nOutputs);
|
|
|
|
|
|
|
|
|
|
for (let i = 0; i < nSamples; i++) {
|
|
|
|
|
for (let j = 0; j < nInputs; j++) {
|
|
|
|
|
featFlat[i * featureDim + j] = features[i][j];
|
|
|
|
|
}
|
|
|
|
|
featFlat[i * featureDim + nInputs] = 1.0;
|
|
|
|
|
for (let j = 0; j < nOutputs; j++) {
|
|
|
|
|
labFlat[i * nOutputs + j] = labels[i][j] || 0;
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
const featPtr = toHeapF32(featFlat);
|
|
|
|
|
const labPtr = toHeapF32(labFlat);
|
2026-04-02 21:35:31 +02:00
|
|
|
const weightPtr = sampleWeights ? toHeapF32(new Float32Array(sampleWeights)) : 0;
|
2026-04-03 18:15:38 +02:00
|
|
|
const lossHistPtr = w.alloc(maxIterations);
|
2026-03-23 22:23:07 +01:00
|
|
|
|
2026-04-03 18:15:38 +02:00
|
|
|
// Train (extended — captures per-iteration loss history)
|
|
|
|
|
const itersRun = w.trainEx(
|
2026-03-23 22:23:07 +01:00
|
|
|
mlp, featPtr, nSamples, featureDim,
|
|
|
|
|
labPtr, nOutputs,
|
2026-04-02 21:35:31 +02:00
|
|
|
weightPtr,
|
2026-04-03 18:15:38 +02:00
|
|
|
learningRate, maxIterations, convergenceThreshold,
|
|
|
|
|
lossHistPtr
|
2026-03-23 22:23:07 +01:00
|
|
|
);
|
|
|
|
|
|
2026-04-03 18:15:38 +02:00
|
|
|
// Read per-iteration loss history
|
|
|
|
|
const lossHistory = fromHeapF32(lossHistPtr, itersRun);
|
|
|
|
|
|
2026-03-23 22:23:07 +01:00
|
|
|
w.free(featPtr);
|
|
|
|
|
w.free(labPtr);
|
2026-04-02 21:35:31 +02:00
|
|
|
if (weightPtr) w.free(weightPtr);
|
2026-04-03 18:15:38 +02:00
|
|
|
w.free(lossHistPtr);
|
|
|
|
|
|
|
|
|
|
const loss = itersRun > 0 ? lossHistory[itersRun - 1] : 0;
|
2026-03-23 22:23:07 +01:00
|
|
|
|
|
|
|
|
// Extract trained weights
|
|
|
|
|
const outPtr = w.alloc(weightCount);
|
|
|
|
|
w.getWeights(mlp, outPtr);
|
|
|
|
|
const trainedWeights = fromHeapF32(outPtr, weightCount);
|
|
|
|
|
w.free(outPtr);
|
|
|
|
|
|
|
|
|
|
self.postMessage({
|
|
|
|
|
type: 'trained',
|
2026-04-03 18:15:38 +02:00
|
|
|
payload: { weights: trainedWeights, loss, lossHistory },
|
2026-03-23 22:23:07 +01:00
|
|
|
});
|
|
|
|
|
}
|
|
|
|
|
};
|