From fb43a4867bfd4dc4ba9c199a958122a748514bb6 Mon Sep 17 00:00:00 2001 From: Tommy Tai Date: Sun, 4 Oct 2026 22:03:31 -0700 Subject: [PATCH 1/2] fix: validate tutorial builds and recover browser runtime errors --- .github/workflows/pages.yml | 52 ++++++++---- materials/README.md | 14 +++- materials/tabicl-explainer/package-lock.json | 77 ++--------------- materials/tabicl-explainer/package.json | 1 + .../website/js/playground-worker-state.js | 34 ++++++++ materials/website/js/playground.js | 8 +- materials/website/js/tabicl/core.js | 82 +++++-------------- materials/website/js/tabicl/tensor.js | 19 ++++- materials/website/js/tabicl/worker.js | 20 +++-- tests/playground-worker-state.test.mjs | 61 ++++++++++++++ tests/tabicl-browser-runtime.test.mjs | 60 ++++++++++++++ tests/tabicl-worker.test.mjs | 66 +++++++++++++++ 12 files changed, 339 insertions(+), 155 deletions(-) create mode 100644 materials/website/js/playground-worker-state.js create mode 100644 tests/playground-worker-state.test.mjs create mode 100644 tests/tabicl-worker.test.mjs diff --git a/.github/workflows/pages.yml b/.github/workflows/pages.yml index 8217d15..f55a0e9 100644 --- a/.github/workflows/pages.yml +++ b/.github/workflows/pages.yml @@ -1,40 +1,40 @@ -name: Deploy TabICL website to GitHub Pages +name: Validate and deploy tutorial website # Publishes materials/website as the repository's GitHub Pages site after # building the standalone Svelte explorer. on: + pull_request: + paths: + - "materials/website/**" + - "materials/tabicl-explainer/**" + - "tests/**" + - ".github/workflows/pages.yml" push: branches: [main] paths: - "materials/website/**" - "materials/tabicl-explainer/**" + - "tests/**" - ".github/workflows/pages.yml" workflow_dispatch: permissions: contents: read - pages: write - id-token: write # Allow one concurrent deployment; don't cancel an in-progress publish. concurrency: - group: pages + group: pages-${{ github.event.pull_request.number || 'publish' }} cancel-in-progress: false jobs: - deploy: + validate: runs-on: ubuntu-latest - environment: - name: github-pages - url: ${{ steps.deployment.outputs.page_url }} + timeout-minutes: 15 steps: - name: Checkout uses: actions/checkout@v4 - - name: Configure Pages - uses: actions/configure-pages@v5 - - name: Set up Node uses: actions/setup-node@v4 with: @@ -42,19 +42,43 @@ jobs: cache: npm cache-dependency-path: materials/tabicl-explainer/package-lock.json + - name: Install explorer dependencies + working-directory: materials/tabicl-explainer + run: npm ci + + - name: Check explorer types + working-directory: materials/tabicl-explainer + run: npm run check + + - name: Test browser runtime + run: node --test tests/*.test.mjs + - name: Build TabICL explorer working-directory: materials/tabicl-explainer - run: | - npm ci - npm run build + run: npm run build - name: Upload website artifact + if: github.event_name != 'pull_request' uses: actions/upload-pages-artifact@v3 with: # Upload only the site folder so it becomes the Pages root; # all asset paths in index.html are relative to this dir. path: materials/website + deploy: + if: github.event_name != 'pull_request' + needs: validate + runs-on: ubuntu-latest + permissions: + pages: write + id-token: write + environment: + name: github-pages + url: ${{ steps.deployment.outputs.page_url }} + steps: + - name: Configure Pages + uses: actions/configure-pages@v5 + - name: Deploy to GitHub Pages id: deployment uses: actions/deploy-pages@v4 diff --git a/materials/README.md b/materials/README.md index 3d985a7..99d6d00 100644 --- a/materials/README.md +++ b/materials/README.md @@ -38,7 +38,8 @@ Use Python 3.11 in a fresh virtual environment: python3.11 -m venv .venv source .venv/bin/activate python -m pip install --require-hashes -r requirements-lock.txt -jupyter lab notebooks/01_tabicl_primer.ipynb +jupyter nbconvert notebooks/01_tabicl_primer.ipynb \ + --to html --execute --ExecutePreprocessor.timeout=2400 ``` `requirements.txt` lists the direct dependencies; `requirements-lock.txt` @@ -48,6 +49,17 @@ notebook verifies its immutable Hugging Face revision and checkpoint checksum before model loading. An internet connection is needed for package, dataset, and checkpoint downloads. +Open `notebooks/01_tabicl_primer.html` to inspect the executed notebook. +The locked environment includes `nbconvert` and the Python kernel. To edit the +notebook in JupyterLab, install the optional editor separately: + +```bash +python -m pip install jupyterlab +jupyter lab notebooks/01_tabicl_primer.ipynb +``` + +JupyterLab is not part of the locked execution environment. + ## Source and attribution - Public tutorial: diff --git a/materials/tabicl-explainer/package-lock.json b/materials/tabicl-explainer/package-lock.json index f0361df..25e8910 100644 --- a/materials/tabicl-explainer/package-lock.json +++ b/materials/tabicl-explainer/package-lock.json @@ -12,6 +12,7 @@ "@sveltejs/kit": "^2.53.3", "@sveltejs/vite-plugin-svelte": "^3.0.0", "@types/d3": "^7.4.3", + "@types/node": "^20.19.43", "autoprefixer": "^10.4.16", "d3": "^7.9.0", "postcss": "^8.4.33", @@ -1664,15 +1665,13 @@ "license": "MIT" }, "node_modules/@types/node": { - "version": "24.5.2", - "resolved": "https://registry.npmjs.org/@types/node/-/node-24.5.2.tgz", - "integrity": "sha512-FYxk1I7wPv3K2XBaoyH2cTnocQEu8AOZ60hPbsyukMPLv5/5qr7V1i8PLHdl6Zf87I+xZXFvPCXYjiTFq+YSDQ==", + "version": "20.19.43", + "resolved": "https://registry.npmjs.org/@types/node/-/node-20.19.43.tgz", + "integrity": "sha512-6oYBAi5ikg4Pl+kGsoYtawUMBT2zZMCvPNF7pVLnHZfd1zf38DRiWn/gT01RYCdUqkv7Fhr+C9ot4/tb+2sVvA==", "dev": true, "license": "MIT", - "optional": true, - "peer": true, "dependencies": { - "undici-types": "~7.12.0" + "undici-types": "~6.21.0" } }, "node_modules/@types/pug": { @@ -3045,18 +3044,6 @@ "node": ">=6" } }, - "node_modules/lilconfig": { - "version": "2.1.0", - "resolved": "https://registry.npmjs.org/lilconfig/-/lilconfig-2.1.0.tgz", - "integrity": "sha512-utWOt/GHzuUxnLKxB6dk81RoOeoNeHgbrXiuGk4yyF5qlRz+iIVWu56E2fqGHFrXz0QNUhLB/8nKqvRH66JKGQ==", - "dev": true, - "license": "MIT", - "optional": true, - "peer": true, - "engines": { - "node": ">=10" - } - }, "node_modules/lines-and-columns": { "version": "1.2.4", "resolved": "https://registry.npmjs.org/lines-and-columns/-/lines-and-columns-1.2.4.tgz", @@ -3487,38 +3474,6 @@ "postcss": "^8.4.21" } }, - "node_modules/postcss-load-config": { - "version": "3.1.4", - "resolved": "https://registry.npmjs.org/postcss-load-config/-/postcss-load-config-3.1.4.tgz", - "integrity": "sha512-6DiM4E7v4coTE4uzA8U//WhtPwyhiim3eyjEMFCnUpzbrkK9wJHgKDT2mR+HbtSrd/NubVaYTOpSpjUl8NQeRg==", - "dev": true, - "license": "MIT", - "optional": true, - "peer": true, - "dependencies": { - "lilconfig": "^2.0.5", - "yaml": "^1.10.2" - }, - "engines": { - "node": ">= 10" - }, - "funding": { - "type": "opencollective", - "url": "https://opencollective.com/postcss/" - }, - "peerDependencies": { - "postcss": ">=8.0.9", - "ts-node": ">=9.0.0" - }, - "peerDependenciesMeta": { - "postcss": { - "optional": true - }, - "ts-node": { - "optional": true - } - } - }, "node_modules/postcss-nested": { "version": "6.2.0", "resolved": "https://registry.npmjs.org/postcss-nested/-/postcss-nested-6.2.0.tgz", @@ -4491,13 +4446,11 @@ } }, "node_modules/undici-types": { - "version": "7.12.0", - "resolved": "https://registry.npmjs.org/undici-types/-/undici-types-7.12.0.tgz", - "integrity": "sha512-goOacqME2GYyOZZfb5Lgtu+1IDmAlAEu5xnD3+xTzS10hT0vzpf0SPjkXwAw9Jm+4n/mQGDP3LO8CPbYROeBfQ==", + "version": "6.21.0", + "resolved": "https://registry.npmjs.org/undici-types/-/undici-types-6.21.0.tgz", + "integrity": "sha512-iwDZqg0QAGrg9Rav5H4n0M64c3mkR59cJ6wQp+7C4nI0gsmExaedaYLNO44eT4AtBBwjbTiGPMlt2Md0T9H9JQ==", "dev": true, - "license": "MIT", - "optional": true, - "peer": true + "license": "MIT" }, "node_modules/update-browserslist-db": { "version": "1.1.3", @@ -4735,18 +4688,6 @@ "integrity": "sha512-l4Sp/DRseor9wL6EvV2+TuQn63dMkPjZ/sp9XkghTEbV9KlPS1xUsZ3u7/IQO4wxtcFB4bgpQPRcR3QCvezPcQ==", "dev": true, "license": "ISC" - }, - "node_modules/yaml": { - "version": "1.10.2", - "resolved": "https://registry.npmjs.org/yaml/-/yaml-1.10.2.tgz", - "integrity": "sha512-r3vXyErRCYJ7wg28yvBY5VSoAF8ZvlcW9/BwUzEtUsjvX/DKs24dIkuwjtuprwJJHsbyUbLApepYTR1BN4uHrg==", - "dev": true, - "license": "ISC", - "optional": true, - "peer": true, - "engines": { - "node": ">= 6" - } } } } diff --git a/materials/tabicl-explainer/package.json b/materials/tabicl-explainer/package.json index 01fbac8..538501e 100644 --- a/materials/tabicl-explainer/package.json +++ b/materials/tabicl-explainer/package.json @@ -13,6 +13,7 @@ "@sveltejs/kit": "^2.53.3", "@sveltejs/vite-plugin-svelte": "^3.0.0", "@types/d3": "^7.4.3", + "@types/node": "^20.19.43", "autoprefixer": "^10.4.16", "d3": "^7.9.0", "postcss": "^8.4.33", diff --git a/materials/website/js/playground-worker-state.js b/materials/website/js/playground-worker-state.js new file mode 100644 index 0000000..ebccd73 --- /dev/null +++ b/materials/website/js/playground-worker-state.js @@ -0,0 +1,34 @@ +export function recoverWorkerError(state, message) { + switch (message.requestType) { + case "inspect": + case "predict": + if (message.tag !== state.queryTag) return false; + state.qInFlight = false; + state.qProba = null; + state.qClass = null; + state.selectedView = null; + state.selectedViewProbability = null; + state.attention = null; + state.attentionBlock = null; + state.attentionHeads = null; + break; + case "prepare": + if (message.tag !== state.prepTag) return false; + state.prepared = false; + state.qInFlight = false; + state.qQueued = false; + break; + case "grid": + if (message.tag !== state.gridTag) return false; + state.gridBusy = false; + state.grid = null; + state.surface = null; + break; + case "load": + state.loading = false; + break; + default: + return false; + } + return true; +} diff --git a/materials/website/js/playground.js b/materials/website/js/playground.js index 6cc3cd1..89ed4d3 100644 --- a/materials/website/js/playground.js +++ b/materials/website/js/playground.js @@ -5,6 +5,7 @@ * Pure vanilla, no build step. */ import { predictKnn } from "./knn.js"; +import { recoverWorkerError } from "./playground-worker-state.js"; import { classIndexFromProbability, contourSegments2d, @@ -745,9 +746,14 @@ function ensureWorker() { updateModelUI(); draw(); } else if (m.type === "error") { - state.loading = false; state.gridBusy = false; updateModelUI(); + if (!recoverWorkerError(state, m)) return; + updateModelUI(); updateReadout(); updateCompanion(); draw(); if (state.model === "tfm") setStatus("error: " + m.message); console.error(m.message); + if ((m.requestType === "inspect" || m.requestType === "predict") && state.qQueued) { + state.qQueued = false; + requestQuery(); + } } }; } diff --git a/materials/website/js/tabicl/core.js b/materials/website/js/tabicl/core.js index 52f3521..1f01b00 100644 --- a/materials/website/js/tabicl/core.js +++ b/materials/website/js/tabicl/core.js @@ -1,4 +1,4 @@ -import { loadFlatTensors, loadManifestAndWeights } from "./tensor.js"; +import { loadFlatTensors, loadManifestAndBuffer, loadManifestAndWeights } from "./tensor.js"; export class UpstreamTensorStore { constructor(manifest, tensors) { @@ -33,8 +33,8 @@ export class CoreBackend { } static async fromBaseUrl(baseUrl, options = {}) { - const tensorStore = await UpstreamTensorStore.load(baseUrl, options); - return new CoreBackend({ task: tensorStore.task, tensorStore }); + const { manifest, buffer } = await loadManifestAndBuffer(baseUrl, options); + return new CoreBackend().loadArtifact(manifest, buffer); } async loadArtifact(manifest, buffer) { @@ -56,42 +56,38 @@ export class CoreBackend { } prepare(XTrain, yTrain, options = {}) { - if (this.model?.prepareContext && options.cacheMode === "repr") { + this.assertModel(); + if (options.cacheMode === "repr") { throw new Error("repr cache mode is not implemented for the JavaScript core yet"); } - if (this.model?.prepareContext) { - return { mode: options.cacheMode || "kv", native: this.model.prepareContext(XTrain, yTrain), options }; - } - return { mode: options.cacheMode || null, XTrain, yTrain, options }; + return { mode: options.cacheMode || "kv", native: this.model.prepareContext(XTrain, yTrain), options }; } predictClassifier(XTrain, yTrain, XTest, nClasses) { - if (this.model?.predict) { - const logits = this.model.predict(XTrain, yTrain, XTest); - return logits.map((row) => row.slice(0, nClasses)); - } - return heuristicClassifier(XTrain, yTrain, XTest, nClasses); + this.assertModel(); + const logits = this.model.predict(XTrain, yTrain, XTest); + return logits.map((row) => row.slice(0, nClasses)); } predictClassifierWithCache(cache, XTest, nClasses) { - if (this.model?.predictQueries && cache?.native) { - return this.model.predictQueries(cache.native, XTest).map((row) => row.slice(0, nClasses)); - } - return this.predictClassifier(cache.XTrain, cache.yTrain, XTest, nClasses); + this.assertModel(); + if (!cache?.native) throw new Error("A fitted model cache is required for prediction"); + return this.model.predictQueries(cache.native, XTest).map((row) => row.slice(0, nClasses)); } predictRegressor(XTrain, yTrain, XTest) { - if (this.model?.predict) { - return this.model.predict(XTrain, yTrain, XTest); - } - return heuristicRegressor(XTrain, yTrain, XTest); + this.assertModel(); + return this.model.predict(XTrain, yTrain, XTest); } predictRegressorWithCache(cache, XTest) { - if (this.model?.predictQueries && cache?.native) { - return this.model.predictQueries(cache.native, XTest); - } - return this.predictRegressor(cache.XTrain, cache.yTrain, XTest); + this.assertModel(); + if (!cache?.native) throw new Error("A fitted model cache is required for prediction"); + return this.model.predictQueries(cache.native, XTest); + } + + assertModel() { + if (!this.model) throw new Error("Load a model artifact before inference"); } } @@ -176,39 +172,3 @@ function translateUpstreamKey(key) { if (match) return `icl_blocks.${match[1]}.${mapTransformerBlockSuffix(match[2])}`; throw new Error(`Unmapped upstream tensor key: ${key}`); } - -function distanceSquared(a, b) { - let acc = 0; - for (let i = 0; i < a.length; i++) { - const delta = Number(a[i]) - Number(b[i]); - acc += delta * delta; - } - return acc; -} - -function heuristicClassifier(XTrain, yTrain, XTest, nClasses) { - return XTest.map((row) => { - const logits = new Array(nClasses).fill(-8); - const weights = new Array(nClasses).fill(0); - for (let i = 0; i < XTrain.length; i++) { - const cls = Number(yTrain[i]); - const w = Math.exp(-distanceSquared(row, XTrain[i])); - if (Number.isInteger(cls) && cls >= 0 && cls < nClasses) weights[cls] += w; - } - for (let cls = 0; cls < nClasses; cls++) logits[cls] = Math.log(weights[cls] + 1e-6); - return logits; - }); -} - -function heuristicRegressor(XTrain, yTrain, XTest) { - return XTest.map((row) => { - let num = 0; - let den = 0; - for (let i = 0; i < XTrain.length; i++) { - const w = Math.exp(-distanceSquared(row, XTrain[i])); - num += w * Number(yTrain[i]); - den += w; - } - return den ? num / den : 0; - }); -} diff --git a/materials/website/js/tabicl/tensor.js b/materials/website/js/tabicl/tensor.js index af4da8c..dcf196a 100644 --- a/materials/website/js/tabicl/tensor.js +++ b/materials/website/js/tabicl/tensor.js @@ -68,11 +68,24 @@ export function loadFlatTensors(manifest, buffer, { dequantize = true } = {}) { return tensors; } -export async function loadManifestAndWeights(baseUrl, { fetchImpl = fetch, dequantize = true } = {}) { - const manifest = await (await fetchImpl(`${baseUrl.replace(/\/$/, "")}/manifest.json`)).json(); +export async function loadManifestAndBuffer(baseUrl, { fetchImpl = fetch } = {}) { + const root = baseUrl.replace(/\/$/, ""); + const manifestResponse = await fetchImpl(`${root}/manifest.json`); + if (!manifestResponse.ok) throw new Error(`Manifest download failed (${manifestResponse.status})`); + const manifest = await manifestResponse.json(); const binaryName = manifest.binary || manifest.bin || "tabicl.bin"; - const buffer = await (await fetchImpl(`${baseUrl.replace(/\/$/, "")}/${binaryName}`)).arrayBuffer(); + const binaryResponse = await fetchImpl(`${root}/${binaryName}`); + if (!binaryResponse.ok) throw new Error(`Model download failed (${binaryResponse.status})`); + const buffer = await binaryResponse.arrayBuffer(); + if (manifest.total_bytes && buffer.byteLength !== manifest.total_bytes) { + throw new Error(`Model byte length ${buffer.byteLength} does not match manifest ${manifest.total_bytes}`); + } await verifySha256(buffer, manifest.binary_sha256, binaryName); + return { manifest, buffer }; +} + +export async function loadManifestAndWeights(baseUrl, { dequantize = true, ...options } = {}) { + const { manifest, buffer } = await loadManifestAndBuffer(baseUrl, options); return { manifest, tensors: loadFlatTensors(manifest, buffer, { dequantize }) }; } diff --git a/materials/website/js/tabicl/worker.js b/materials/website/js/tabicl/worker.js index 4f2fa55..d626392 100644 --- a/materials/website/js/tabicl/worker.js +++ b/materials/website/js/tabicl/worker.js @@ -6,7 +6,7 @@ * in {type:'inspect', X, tag, viewIndex} -> {type:'inspection', tag, proba, selectedViewProbability, selectedView, selectedViewTrace, attention, block, heads} * in {type:'grid', X, tag, chunk} -> repeated {type:'gridChunk', tag, start, proba} then {type:'gridDone', tag} * in {type:'cancelGrid'} -> stop the current grid between chunks - * any failure -> {type:'error', message} + * any failure -> {type:'error', requestType, tag, message} */ import { TabICLClassifier } from "./classifier.js"; import { CoreBackend } from "./core.js"; @@ -20,11 +20,14 @@ self.onmessage = async (e) => { if (m.type === "load") { if (backend) { self.postMessage({ type: "loaded", bytes: 0 }); return; } const base = new URL(m.modelBase || "../../model/", import.meta.url); - const manifest = await (await fetch(new URL("manifest.json", base), { cache: "no-store" })).json(); + const manifestResponse = await fetch(new URL("manifest.json", base), { cache: "no-store" }); + if (!manifestResponse.ok) throw new Error(`Manifest download failed (${manifestResponse.status})`); + const manifest = await manifestResponse.json(); if (manifest.schema !== "tabicl-browser-js/flat-tensors-v1") { throw new Error("Expected the TabICLv2 browser manifest; clear the stale site cache and reload"); } const resp = await fetch(new URL(manifest.binary || "tabicl.bin", base), { cache: "force-cache" }); + if (!resp.ok) throw new Error(`Model download failed (${resp.status})`); const total = +(resp.headers.get("Content-Length") || 0); const reader = resp.body.getReader(); const chunks = []; let loaded = 0; @@ -40,13 +43,15 @@ self.onmessage = async (e) => { throw new Error(`Model byte length ${loaded} does not match manifest ${manifest.total_bytes}`); } await verifySha256(buf.buffer, manifest.binary_sha256, manifest.binary || "tabicl.bin"); - backend = new CoreBackend({ task: "classifier" }); - await backend.loadClassifierArtifact(manifest, buf.buffer); + const loadedBackend = new CoreBackend({ task: "classifier" }); + await loadedBackend.loadClassifierArtifact(manifest, buf.buffer); + backend = loadedBackend; self.postMessage({ type: "loaded", bytes: loaded }); } else if (m.type === "prepare") { activeGridTag = null; + estimator = null; if (!backend) throw new Error("Load the browser classifier before preparing context"); - estimator = new TabICLClassifier({ + const preparedEstimator = new TabICLClassifier({ nEstimators: 8, normMethods: null, featShuffleMethod: "latin", @@ -59,7 +64,8 @@ self.onmessage = async (e) => { randomState: 42, backend, }); - estimator.fit(m.X, m.y, { kvCache: true }); + preparedEstimator.fit(m.X, m.y, { kvCache: true }); + estimator = preparedEstimator; self.postMessage({ type: "prepared", tag: m.tag, viewCount: 8 }); } else if (m.type === "predict") { if (!estimator) throw new Error("Prepare the classifier before prediction"); @@ -103,6 +109,6 @@ self.onmessage = async (e) => { activeGridTag = null; } } catch (err) { - self.postMessage({ type: "error", message: String((err && err.stack) || err) }); + self.postMessage({ type: "error", requestType: m.type, tag: m.tag, message: String((err && err.stack) || err) }); } }; diff --git a/tests/playground-worker-state.test.mjs b/tests/playground-worker-state.test.mjs new file mode 100644 index 0000000..c698f1b --- /dev/null +++ b/tests/playground-worker-state.test.mjs @@ -0,0 +1,61 @@ +import assert from "node:assert/strict"; +import test from "node:test"; + +import { recoverWorkerError } from "../materials/website/js/playground-worker-state.js"; + +function pendingState() { + return { + loading: true, prepared: true, prepTag: 3, + qInFlight: true, qQueued: true, queryTag: 7, + qProba: [0.3, 0.7], qClass: 1, + selectedView: { viewIndex: 0 }, selectedViewProbability: [0.2, 0.8], + attention: [1], attentionBlock: 12, attentionHeads: 8, + gridBusy: true, gridTag: 5, grid: new Float32Array([0.4]), surface: {}, + }; +} + +test("a failed query releases the request and preserves queued work and the active grid", () => { + const state = pendingState(); + assert.equal(recoverWorkerError(state, { requestType: "inspect", tag: 7 }), true); + assert.equal(state.qInFlight, false); + assert.equal(state.qQueued, true); + assert.equal(state.qProba, null); + assert.equal(state.selectedView, null); + assert.equal(state.attention, null); + assert.equal(state.loading, true); + assert.equal(state.gridBusy, true); + assert.equal(state.prepared, true); +}); + +test("errors from an old query, context, or grid do not clear current requests", () => { + for (const requestType of ["inspect", "predict", "prepare", "grid"]) { + const state = pendingState(); + const before = structuredClone(state); + assert.equal(recoverWorkerError(state, { requestType, tag: -1 }), false); + assert.deepEqual(state, before); + } +}); + +test("a failed grid clears the partial field without releasing an active query", () => { + const state = pendingState(); + assert.equal(recoverWorkerError(state, { requestType: "grid", tag: 5 }), true); + assert.equal(state.gridBusy, false); + assert.equal(state.grid, null); + assert.equal(state.surface, null); + assert.equal(state.qInFlight, true); +}); + +test("a failed context prevents prediction until preparation succeeds", () => { + const state = pendingState(); + assert.equal(recoverWorkerError(state, { requestType: "prepare", tag: 3 }), true); + assert.equal(state.prepared, false); + assert.equal(state.qInFlight, false); + assert.equal(state.qQueued, false); +}); + +test("a failed load permits retry without clearing unrelated work", () => { + const state = pendingState(); + assert.equal(recoverWorkerError(state, { requestType: "load" }), true); + assert.equal(state.loading, false); + assert.equal(state.qInFlight, true); +}); diff --git a/tests/tabicl-browser-runtime.test.mjs b/tests/tabicl-browser-runtime.test.mjs index a7773c2..a08f854 100644 --- a/tests/tabicl-browser-runtime.test.mjs +++ b/tests/tabicl-browser-runtime.test.mjs @@ -27,6 +27,66 @@ const X = [ const y = [0, 0, 0, 0, 1, 1, 1, 1]; const query = [[0.25, -0.15, 0.35]]; +test("an unloaded backend rejects inference instead of returning heuristic predictions", () => { + const backend = new CoreBackend(); + for (const predict of [ + () => backend.prepare(X, y), + () => backend.predictClassifier(X, y, query, 2), + () => backend.predictClassifierWithCache({ XTrain: X, yTrain: y }, query, 2), + () => backend.predictRegressor(X, y, query), + () => backend.predictRegressorWithCache({ XTrain: X, yTrain: y }, query), + () => new TabICLClassifier().fit(X, y).predictProba(query), + ]) { + assert.throws(predict, /Load a model artifact before inference/); + } +}); + +test("loading a backend from a URL creates an executable artifact model", async () => { + const manifest = JSON.parse(readFileSync(manifestPath, "utf8")); + const bytes = readFileSync(binaryPath); + const requests = []; + const backend = await CoreBackend.fromBaseUrl("https://example.test/model/", { + fetchImpl: async (url) => { + requests.push(url); + if (url.endsWith("/manifest.json")) return Response.json(manifest); + if (url.endsWith("/tabicl.bin")) return new Response(bytes); + throw new Error(`Unexpected model request: ${url}`); + }, + }); + + assert.deepEqual(requests, [ + "https://example.test/model/manifest.json", + "https://example.test/model/tabicl.bin", + ]); + const cache = backend.prepare(X, y); + const logits = backend.predictClassifierWithCache(cache, query, 2); + assert.equal(logits.length, 1); + assert.equal(logits[0].length, 2); + assert.ok(logits[0].every(Number.isFinite)); + assert.throws(() => backend.predictClassifierWithCache({}, query, 2), /fitted model cache/); +}); + +test("backend downloads reject HTTP errors and altered model bytes", async () => { + await assert.rejects(CoreBackend.fromBaseUrl("https://example.test/model", { + fetchImpl: async () => new Response("Not found", { status: 404 }), + }), /Manifest download failed \(404\)/); + + const manifest = JSON.parse(readFileSync(manifestPath, "utf8")); + await assert.rejects(CoreBackend.fromBaseUrl("https://example.test/model", { + fetchImpl: async (url) => url.endsWith("/manifest.json") + ? Response.json(manifest) + : new Response("Unavailable", { status: 503 }), + }), /Model download failed \(503\)/); + + const alteredBytes = Buffer.from(readFileSync(binaryPath)); + alteredBytes[0] ^= 1; + await assert.rejects(CoreBackend.fromBaseUrl("https://example.test/model", { + fetchImpl: async (url) => url.endsWith("/manifest.json") + ? Response.json(manifest) + : new Response(alteredBytes), + }), /SHA-256 mismatch/); +}); + test("browser manifest translation is idempotent for stale mapped manifests", () => { const mapped = toNanoClassifierManifest({ config: { embed_dim: 128, n_cls_cols: 4, feature_group_size: 3 }, diff --git a/tests/tabicl-worker.test.mjs b/tests/tabicl-worker.test.mjs new file mode 100644 index 0000000..dc31ce8 --- /dev/null +++ b/tests/tabicl-worker.test.mjs @@ -0,0 +1,66 @@ +import assert from "node:assert/strict"; +import { once } from "node:events"; +import test from "node:test"; +import { Worker } from "node:worker_threads"; + +const workerUrl = new URL("../materials/website/js/tabicl/worker.js", import.meta.url).href; + +function browserWorker(t, setup = "") { + const worker = new Worker(` + const { parentPort } = require("node:worker_threads"); + globalThis.self = { postMessage: (message) => parentPort.postMessage(message) }; + ${setup} + import(${JSON.stringify(workerUrl)}).then(() => { + parentPort.on("message", (data) => self.onmessage({ data })); + }).catch((error) => { throw error; }); + `, { eval: true }); + t.after(() => worker.terminate()); + return worker; +} + +async function send(worker, message) { + const response = once(worker, "message"); + worker.postMessage(message); + const [result] = await response; + return result; +} + +test("worker failures retain the request type and tag", { timeout: 10_000 }, async (t) => { + const worker = browserWorker(t); + for (const type of ["prepare", "predict", "inspect", "grid"]) { + const response = await send(worker, { type, tag: 123, X: [[0]], y: [0] }); + assert.equal(response.type, "error"); + assert.equal(response.requestType, type); + assert.equal(response.tag, 123); + assert.match(response.message, /Load|Prepare/); + } +}); + +test("a failed artifact initialization cannot be reported as loaded on retry", { timeout: 10_000 }, async (t) => { + const worker = browserWorker(t, ` + const manifest = { + schema: "tabicl-browser-js/flat-tensors-v1", task: "classifier", config: {}, + binary_sha256: "e3b0c44298fc1c149afbf4c8996fb92427ae41e4649b934ca495991b7852b855", + tensors: [{ name: "unknown_tensor", shape: [1] }], + }; + globalThis.fetch = async (url) => String(url).endsWith("manifest.json") + ? Response.json(manifest) : new Response(new Uint8Array()); + `); + for (const tag of [1, 2]) { + const response = await send(worker, { type: "load", tag }); + assert.equal(response.type, "error"); + assert.equal(response.requestType, "load"); + assert.equal(response.tag, tag); + assert.match(response.message, /Unmapped upstream tensor key/); + } +}); + +test("worker reports a failed manifest download before reading its body", { timeout: 10_000 }, async (t) => { + const worker = browserWorker(t, ` + globalThis.fetch = async () => new Response("Not found", { status: 404 }); + `); + const response = await send(worker, { type: "load", tag: 9 }); + assert.equal(response.type, "error"); + assert.equal(response.tag, 9); + assert.match(response.message, /Manifest download failed \(404\)/); +}); From 55f8c28ac0875216d83254ba418ce6d6c0aaa856 Mon Sep 17 00:00:00 2001 From: Tommy Tai Date: Sun, 4 Oct 2026 23:18:37 -0700 Subject: [PATCH 2/2] docs: add macOS notebook prerequisites --- materials/README.md | 9 +++++++++ 1 file changed, 9 insertions(+) diff --git a/materials/README.md b/materials/README.md index 99d6d00..5abc4f6 100644 --- a/materials/README.md +++ b/materials/README.md @@ -32,6 +32,15 @@ or Apple MPS on supported local hardware, with CPU as the fallback. ### Local or CI +On macOS, XGBoost also needs the OpenMP runtime: + +```bash +brew install libomp +``` + +Use Python and OpenMP for the same processor architecture. On Apple Silicon, +use native `arm64` Python and Homebrew. + Use Python 3.11 in a fresh virtual environment: ```bash