diff --git a/.github/workflows/e2e.yml b/.github/workflows/e2e.yml index fc4b3cc..f7c555f 100644 --- a/.github/workflows/e2e.yml +++ b/.github/workflows/e2e.yml @@ -1,14 +1,33 @@ name: E2E Tests +# MANUAL ONLY. These do not run on pull requests or pushes. +# +# The team's workflow is to run e2e locally before pushing (`npm run test:e2e`), which made this +# redundant on every PR — and expensive: 13 minutes, of which 84% was the tests themselves. Playwright +# runs with `workers: 1` in CI, deliberately, because parallel workers contend for the single dev +# server and reintroduce flake. So the cost was not fixable by caching or a faster runner; only by +# sharding across jobs, which is more machinery than a locally-run suite justifies. +# +# The workflow is kept rather than deleted, and can be started from the Actions tab or with +# `gh workflow run e2e.yml`. Worth doing before a release, or on any change to what the app +# DOWNLOADS: the byte-budget guards in 18-meme-hit-likelihood.spec.js and 19-axomeme.spec.js exist +# because a previous version shipped 13.5 MB of ML runtime to every method and rendered nothing, and +# an "is the element absent?" assertion passed the whole time. Those are the tests least likely to be +# run from memory and most likely to catch a regression nobody thought to look for. on: - pull_request: - branches: [main] - push: - branches: [main] + workflow_dispatch: + inputs: + suite: + description: 'Which suite to run' + required: false + default: 'fast' + type: choice + options: [fast, wasm, both] jobs: fast-e2e: name: Fast E2E Tests + if: inputs.suite == 'fast' || inputs.suite == 'both' runs-on: ubuntu-latest steps: - uses: actions/checkout@v4 @@ -40,8 +59,10 @@ jobs: wasm-e2e: name: WASM E2E Tests + # Previously main-only. With no push trigger left, that condition would never fire, so the + # selector above is what decides now. + if: inputs.suite == 'wasm' || inputs.suite == 'both' runs-on: ubuntu-latest - if: github.ref == 'refs/heads/main' steps: - uses: actions/checkout@v4 diff --git a/.gitignore b/.gitignore index 3d702c1..a7ccfab 100644 --- a/.gitignore +++ b/.gitignore @@ -74,3 +74,7 @@ coverage-e2e /benchmark-results/ /hyphy-wasm-artifact/ /src/benchmark/test-alignments/ + +# onnxruntime-web WASM runtime, copied from node_modules at build time by +# scripts/copy-ort-wasm.mjs. Pinned by package.json, not by git. +static/ort/ diff --git a/e2e/18-meme-hit-likelihood.spec.js b/e2e/18-meme-hit-likelihood.spec.js index 5c59969..53ec065 100644 --- a/e2e/18-meme-hit-likelihood.spec.js +++ b/e2e/18-meme-hit-likelihood.spec.js @@ -6,7 +6,9 @@ * 2. WHAT IT COSTS. The original version of this feature downloaded 13.5 MB of ML-runtime WASM * for every method, including the ~14 that have no estimate, and then rendered nothing. An * "is the element absent?" assertion passed the whole time, because the element WAS absent — - * the bytes were the bug. So these tests count bytes, not just elements. + * the bytes were the bug. So these tests count bytes, not just elements. DM3 now ships + * onnxruntime-web for AxoMEME, which makes that regression cheap to reintroduce by accident; + * see RUNTIME_TOMBSTONE for why these flows must still transfer none of it. * * WHAT MOVED SINCE THE LAST REVISION. The estimator is no longer a converted coefficient table; it * is XGBoost's own `save_model()` JSON, shipped verbatim and walked in ~40 lines of plain JS. That @@ -159,12 +161,25 @@ const ESTIMATOR_MARKS = ['meme-hit-likelihood', 'hit-basis-toggle']; const SCOPE_LEAF = /prescreen\/scope\.js/; /** Dev-server source paths. Kept as a belt-and-braces net alongside the content markers. */ -const ESTIMATOR_ASSET = /prescreen|hit_likelihood|hitLikelihood|meme_gate|onnxruntime|ort-wasm|tfjs/i; +const ESTIMATOR_ASSET = + /prescreen|hit_likelihood|hitLikelihood|meme_gate|onnxruntime|ort-wasm|tfjs/i; /** - * A TOMBSTONE, not a live dependency. These names belong to the ML runtimes DM3 does not ship and - * must never start shipping again; the first version of this feature pulled 13.5 MB of the first - * one. Nothing else in the repository references them — that is the state this defends. + * A TOMBSTONE — and READ THIS BEFORE RELAXING IT, because the sentence that used to justify it is + * no longer true while the assertions themselves are unchanged and still exactly right. + * + * It used to say "nothing else in the repository references these names — that is the state this + * defends". DM3 now has a legitimate reason to depend on onnxruntime-web: AxoMEME 2.0 is a 3.78 MB + * transformer whose graph is a real neural network, and it cannot be walked in plain JS the way the + * gate's 500 three-feature trees can. So the runtime IS in package.json now. + * + * That does NOT weaken anything below, because these tests are scoped to two flows — select FEL, + * and select MEME — and neither one is AxoMEME. The invariant they defend was always per-flow, not + * per-repository: A METHOD MUST NOT PAY FOR A MODEL IT DOES NOT RENDER. The first version of this + * feature charged all fifteen methods 13.5 MB for an estimate fourteen of them never showed. Now + * that a runtime is a real dependency that resolves instead of erroring, an accidental static + * import would silently pull ~13 MB into the main graph and every one of these flows would pay it. + * The guard therefore matters MORE than it did when it was written, not less. * * Checked against EVERY response, and against response BODIES as well as URLs. Previously the byte * recorder discarded anything that did not already match ESTIMATOR_ASSET before the tombstone was @@ -333,7 +348,6 @@ test.describe('MEME hit-likelihood', () => { await expect(panel.locator('[data-testid="meme-hit-likelihood"]')).toHaveCount(1); await expect(panel.locator('.timing-estimate')).toHaveCount(1); }); - }); /** @@ -348,6 +362,46 @@ test.describe('MEME hit-likelihood', () => { * less methods pay for it again), was fetched before the listener existed and was therefore * invisible. Each test below owns its whole page lifecycle, with the recorder attached first. */ +test.describe('MEME hit-likelihood layout', () => { + test.setTimeout(180000); + + test('the Run button does not move when the estimate resolves', async ({ page }) => { + // The panel sits directly above the Run button, so if the pending placeholder is not exactly as + // tall as the mounted row, the button shifts under a pointer already travelling towards it. + // That reserve used to be two hand-measured pixel constants in RunOutlook, duplicated from + // MemeHitLikelihood.css with a comment telling the next person to re-measure — and the comment + // recorded that they were already 5.5px stale. They are now derived with calc() from tokens + // declared once in app.css, and this asserts the property those numbers existed to protect. + await freshStart(page); + await loadDemoFile(page, 'CD2-slim.fna'); + await expect(async () => { + await goToAnalyzeTab(page); + await expect(page.locator('[data-testid="method-dropdown"]')).toBeVisible({ timeout: 5000 }); + }).toPass({ timeout: 60000 }); + + // The tokens must resolve from the always-loaded stylesheet, not the lazy chunk — the + // placeholder is drawn before that chunk exists. + const reserve = await page.evaluate(() => + getComputedStyle(document.documentElement).getPropertyValue('--hit-body-reserve').trim() + ); + expect(reserve, '--hit-body-reserve is not defined in the always-loaded stylesheet').toMatch( + /\d/ + ); + + await selectMethod(page, 'MEME'); + const runButton = page.locator('[data-testid="run-analysis-btn"]'); + await runButton.waitFor({ timeout: 30000 }); + const before = (await runButton.boundingBox())?.y ?? 0; + // Long enough for the estimator chunk, the model fetch and the score to all land. + await page.waitForTimeout(9000); + const after = (await runButton.boundingBox())?.y ?? 0; + expect( + Math.abs(after - before), + 'the Run button moved when the estimate resolved' + ).toBeLessThan(2); + }); +}); + test.describe('MEME hit-likelihood cost', () => { test.setTimeout(180000); @@ -455,8 +509,10 @@ test.describe('MEME hit-likelihood cost', () => { // And the row really is driven by the estimator chunk, not by something inlined into the // main bundle — otherwise "the model is fetched lazily" would be true and irrelevant. - expect(all.filter(isEstimatorCode).length, 'the estimator chunk was never fetched') - .toBeGreaterThan(0); + expect( + all.filter(isEstimatorCode).length, + 'the estimator chunk was never fetched' + ).toBeGreaterThan(0); // No ML runtime, now or ever again — checked over every response in the session, by URL and // by content. @@ -506,11 +562,7 @@ test.describe('MEME hit-likelihood routing', () => { test('a discouraging estimate says what about the data would change it', async ({ page }) => { // 4 taxa x 200 codons, two substitutions each on disjoint sites: the thin-and-shallow shape // that the low band is made of (its real population medians 4 sequences and ~140 codons). - await uploadAndAnalyze( - page, - 'synthetic-thin.fna', - lowDivergenceAlignment(4, 200, 2) - ); + await uploadAndAnalyze(page, 'synthetic-thin.fna', lowDivergenceAlignment(4, 200, 2)); await selectMethod(page, 'MEME'); const row = page.locator('[data-testid="meme-hit-likelihood"]'); diff --git a/e2e/19-axomeme.spec.js b/e2e/19-axomeme.spec.js new file mode 100644 index 0000000..93e0a94 --- /dev/null +++ b/e2e/19-axomeme.spec.js @@ -0,0 +1,260 @@ +/** + * E2E tests for AxoMEME, the in-browser neural surrogate for MEME. + * + * These exist because AxoMEME's failure modes are ones no unit test can see. The model and the + * runtime are fetched at run time from static assets, and every bug found while wiring this up was + * in that fetch rather than in any computation: + * + * 1. onnxruntime-web does not bundle its WASM binary. With no `wasmPaths` it resolves to a jsDelivr + * CDN — which violates this project's core constraint that the site be servable with no other + * domains involved, and which simply fails offline. The runtime reports "no available backend + * found", which reads like a broken model rather than a missing asset. + * 2. The DEFAULT package entry wants the JSEP (WebGPU) binary, a different 26.8 MB file from the + * CPU one. Importing `onnxruntime-web` instead of `onnxruntime-web/wasm` fails with the same + * opaque message even when the CPU binary is present and served correctly. + * + * Both were found by hand and neither would have failed a unit test. So what these tests assert is + * WHICH URLS THE PAGE ACTUALLY REQUESTS, not just that a number came out. + */ + +import { test, expect } from './fixtures/coverage.js'; +import { + freshStart, + loadDemoFile, + goToAnalyzeTab, + selectMethod, + clickRunAnalysis +} from './fixtures/helpers.js'; + +/** Anything served from another origin. The project must never need one. */ +const isThirdParty = (url) => !/^https?:\/\/(localhost|127\.0\.0\.1)/.test(url); + +/** The CPU runtime this feature is built against. */ +const CPU_WASM = /ort-wasm-simd-threaded\.(wasm|mjs)$/; +/** The WebGPU / asyncify / JSPI variants, none of which should ever be requested. */ +const OTHER_VARIANTS = /ort-wasm-simd-threaded\.(jsep|asyncify|jspi)\./; + +function trackRequests(page) { + const urls = []; + page.on('request', (r) => urls.push(r.url())); + return () => urls; +} + +test.describe('AxoMEME', () => { + test.setTimeout(180000); + + test('appears in the method list and describes itself as a prediction', async ({ page }) => { + await freshStart(page); + await loadDemoFile(page, 'CD2-slim.fna'); + await expect(async () => { + await goToAnalyzeTab(page); + await expect(page.locator('[data-testid="method-dropdown"]')).toBeVisible({ timeout: 5000 }); + }).toPass({ timeout: 60000 }); + + await selectMethod(page, 'AxoMEME'); + + // The copy has to say PREDICT. AxoMEME estimates what MEME would report; presenting it as a + // completed selection analysis is the one framing error that matters here, and the method + // description is where a user forms that expectation. + const body = await page.locator('body').innerText(); + expect(body).toMatch(/predict/i); + }); + + test('fetches its runtime and model from THIS origin, never a CDN', async ({ page }) => { + // The project constraint, asserted rather than trusted: "The website needs to be entirely + // self-contained so that it can be served locally. No pulling down from other domains." + const requests = trackRequests(page); + await freshStart(page); + await loadDemoFile(page, 'CD2-slim.fna'); + await expect(async () => { + await goToAnalyzeTab(page); + await expect(page.locator('[data-testid="method-dropdown"]')).toBeVisible({ timeout: 5000 }); + }).toPass({ timeout: 60000 }); + await selectMethod(page, 'AxoMEME'); + + expect(await clickRunAnalysis(page), 'run button was not clickable').toBe(true); + + // Give the runtime time to fetch its binary and score 17 sites. + await page.waitForTimeout(45000); + + const all = requests(); + // SCOPED TO AXOMEME'S OWN ASSETS, deliberately. A broader filter here fails, and not because of + // anything in this feature: the app already fetches katex from cdn.jsdelivr.net and aioli from + // biowasm.com, which are real violations of the same project constraint and predate this work. + // Widening this assertion would make an AxoMEME test fail for someone else's reason, which is + // how a guard gets deleted. They are logged separately instead. + const offOrigin = all.filter(isThirdParty).filter((u) => /onnx|ort-wasm/i.test(u)); + expect(offOrigin, `AxoMEME assets fetched off-origin: ${offOrigin.join(', ')}`).toEqual([]); + + // And it must have asked for the CPU binary specifically. + const wasmRequests = all.filter((u) => CPU_WASM.test(u)); + expect(wasmRequests.length, 'the CPU WASM runtime was never requested').toBeGreaterThan(0); + for (const u of wasmRequests) expect(u).toMatch(/\/ort\//); + + // The JSEP / asyncify / JSPI builds are 15-27 MB each and this feature uses none of them. + // Requesting one means the default package entry crept back in. + const wrongVariant = all.filter((u) => OTHER_VARIANTS.test(u)); + expect( + wrongVariant, + `requested a runtime variant we do not ship: ${wrongVariant.join(', ')}` + ).toEqual([]); + }); + + test('produces a per-site result table', async ({ page }) => { + await freshStart(page); + await loadDemoFile(page, 'CD2-slim.fna'); + await expect(async () => { + await goToAnalyzeTab(page); + await expect(page.locator('[data-testid="method-dropdown"]')).toBeVisible({ timeout: 5000 }); + }).toPass({ timeout: 60000 }); + await selectMethod(page, 'AxoMEME'); + + expect(await clickRunAnalysis(page), 'run button was not clickable').toBe(true); + + // The analysis must reach a terminal state, and it must be success. An error toast here is the + // signal that the runtime or the model failed to load. + await expect(page.getByText(/AxoMEME prediction complete/i)).toBeVisible({ timeout: 120000 }); + + // And the results must actually render. A completed analysis with nothing to show is the + // failure mode this test exists to catch. + await page + .getByRole('button', { name: /results/i }) + .first() + .click(); + // The Results tab lists every analysis and defaults the detail pane to the FIRST one, which is + // the datareader job every upload creates. Selecting the AxoMEME run explicitly is required; + // without it this test passes or fails on whatever happens to be first. + await page + .locator('.analysis-card') + .filter({ hasText: /AXOMEME Analysis/i }) + .getByRole('button', { name: 'View' }) + .first() + .click(); + await expect(page.getByText(/AxoMEME predictions/i)).toBeVisible({ timeout: 30000 }); + + // The framing is not decoration. A researcher reading a per-site table is one step from + // writing it up, so the page must say MEME was not run AND that the number is not a p-value. + const body = await page.locator('body').innerText(); + expect(body).toMatch(/MEME was not run/i); + expect(body).toMatch(/not calibrated/i); + + // The table and plots come from hyphy-scope's AxomemeVisualization; DataMonkey contributes only + // the caveats about what it did to the data first. These assertions are on the library's + // output, so they also catch a stale `npm link` or a package build that did not include it. + await expect(page.locator('.axomeme-visualization')).toBeVisible(); + await expect(page.locator('.axomeme-plot svg').first()).toBeVisible({ timeout: 20000 }); + + // Rank is a first-class column and the score is NOT labelled LRT. Measured across 12 real + // submissions, the model's predicted LRT clears the chi-square gates once in 662 sites, so + // presenting it as an LRT invites a comparison it cannot support. + await expect(page.getByRole('columnheader', { name: 'Percentile' })).toBeVisible(); + // exact: log(1+score) also matches a loose "Score". + await expect(page.getByRole('columnheader', { name: 'Score', exact: true })).toBeVisible(); + await expect(page.getByRole('columnheader', { name: 'LRT' })).toHaveCount(0); + await expect(page.getByRole('columnheader', { name: /dN/ })).toBeVisible(); + }); + + test('is marked Beta wherever a user meets it', async ({ page }) => { + // The MODEL is under active development, not the integration. The badge has to appear both + // where the method is chosen and where the numbers are read — someone landing on a shared + // results view never sees the selector. + await freshStart(page); + await loadDemoFile(page, 'CD2-slim.fna'); + await expect(async () => { + await goToAnalyzeTab(page); + await expect(page.locator('[data-testid="method-dropdown"]')).toBeVisible({ timeout: 5000 }); + }).toPass({ timeout: 60000 }); + + await selectMethod(page, 'AxoMEME'); + await expect(page.locator('.beta-badge')).toBeVisible({ timeout: 10000 }); + + // And it must be specific to AxoMEME rather than decorating every method. + await selectMethod(page, 'FEL'); + await expect(page.locator('.beta-badge')).toHaveCount(0); + + await selectMethod(page, 'AxoMEME'); + expect(await clickRunAnalysis(page), 'run button was not clickable').toBe(true); + await expect(page.getByText(/AxoMEME prediction complete/i)).toBeVisible({ timeout: 120000 }); + await page + .getByRole('button', { name: /results/i }) + .first() + .click(); + await page + .locator('.analysis-card') + .filter({ hasText: /AXOMEME Analysis/i }) + .getByRole('button', { name: 'View' }) + .first() + .click(); + await expect(page.locator('.axomeme-beta')).toBeVisible({ timeout: 30000 }); + }); + + test('suppresses controls that cannot reach the model', async ({ page }) => { + // Both were live at first and both imply a capability AxoMEME does not have: there is no + // server-side AxoMEME, and the model's tokenizer bakes in the universal genetic code. + await freshStart(page); + await loadDemoFile(page, 'CD2-slim.fna'); + await expect(async () => { + await goToAnalyzeTab(page); + await expect(page.locator('[data-testid="method-dropdown"]')).toBeVisible({ timeout: 5000 }); + }).toPass({ timeout: 60000 }); + + await selectMethod(page, 'AxoMEME'); + const body = await page.locator('body').innerText(); + expect(body).not.toMatch(/Backend Server/i); + expect(body).toMatch(/Runs in your browser/i); + expect(body).toMatch(/cannot be changed for this method/i); + + // And the controls must come back for a method that does have them. + await selectMethod(page, 'FEL'); + const felBody = await page.locator('body').innerText(); + expect(felBody).toMatch(/Backend Server/i); + expect(felBody).toMatch(/Genetic Code/i); + }); +}); + +test.describe('AxoMEME on the bundled demos', () => { + test.setTimeout(300000); + + // large.nex is the regression that prompted this block. Its tree carries negative branch lengths + // from DM3's own NJ inference — a patristic sum of -1.04e-5, which is zero with rounding error on + // it — and the input check rejected the whole run for it. Every bundled demo now runs, so a + // tightened validator cannot quietly break the ones a user is most likely to click. + for (const demo of ['small.nex', 'medium.nex', 'large.nex']) { + test(`runs on ${demo}`, async ({ page }) => { + await freshStart(page); + await loadDemoFile(page, demo); + await expect(async () => { + await goToAnalyzeTab(page); + await expect(page.locator('[data-testid="method-dropdown"]')).toBeVisible({ + timeout: 5000 + }); + }).toPass({ timeout: 60000 }); + await selectMethod(page, 'AxoMEME'); + expect(await clickRunAnalysis(page), 'run button was not clickable').toBe(true); + await expect(page.getByText(/AxoMEME prediction complete/i)).toBeVisible({ timeout: 180000 }); + }); + } +}); + +test.describe('AxoMEME cost', () => { + test.setTimeout(180000); + + test('a non-AxoMEME method downloads none of the runtime or the model', async ({ page }) => { + // The invariant the whole session.js design exists for: AxoMEME is one method of fifteen, and + // the other fourteen must not pay 17 MB for it. This is the same guard the MEME hit-likelihood + // suite applies to its own model, and it matters more here because the payload is larger. + const requests = trackRequests(page); + await freshStart(page); + await loadDemoFile(page, 'CD2-slim.fna'); + await expect(async () => { + await goToAnalyzeTab(page); + await expect(page.locator('[data-testid="method-dropdown"]')).toBeVisible({ timeout: 5000 }); + }).toPass({ timeout: 60000 }); + + await selectMethod(page, 'FEL'); + await page.waitForTimeout(3000); + + const offending = requests().filter((u) => /ort-wasm|onnxruntime|axomeme.*\.onnx/i.test(u)); + expect(offending, `FEL pulled AxoMEME assets: ${offending.join(', ')}`).toEqual([]); + }); +}); diff --git a/package-lock.json b/package-lock.json index 16766a7..12443c9 100644 --- a/package-lock.json +++ b/package-lock.json @@ -13,10 +13,11 @@ "@sentry/sveltekit": "^9", "alivibe": ">=0.1.0", "d3": "^7.9.0", - "hyphy-scope": "^1.9.1", + "hyphy-scope": "^1.10.0", "jszip": "^3.10.1", "lucide-svelte": "^0.559.0", "marked": "^15.0.12", + "onnxruntime-web": "^1.27.0", "phylotree": "^2.2.1", "socket.io-client": "^4.8.1", "socket.io-stream": "^0.9.1", @@ -2677,6 +2678,63 @@ "@opentelemetry/api": "^1.8" } }, + "node_modules/@protobufjs/aspromise": { + "version": "1.1.2", + "resolved": "https://registry.npmjs.org/@protobufjs/aspromise/-/aspromise-1.1.2.tgz", + "integrity": "sha512-j+gKExEuLmKwvz3OgROXtrJ2UG2x8Ch2YZUxahh+s1F2HZ+wAceUNLkvy6zKCPVRkU++ZWQrdxsUeQXmcg4uoQ==", + "license": "BSD-3-Clause" + }, + "node_modules/@protobufjs/base64": { + "version": "1.1.2", + "resolved": "https://registry.npmjs.org/@protobufjs/base64/-/base64-1.1.2.tgz", + "integrity": "sha512-AZkcAA5vnN/v4PDqKyMR5lx7hZttPDgClv83E//FMNhR2TMcLUhfRUBHCmSl0oi9zMgDDqRUJkSxO3wm85+XLg==", + "license": "BSD-3-Clause" + }, + "node_modules/@protobufjs/codegen": { + "version": "2.0.5", + "resolved": "https://registry.npmjs.org/@protobufjs/codegen/-/codegen-2.0.5.tgz", + "integrity": "sha512-zgXFLzW3Ap33e6d0Wlj4MGIm6Ce8O89n/apUaGNB/jx+hw+ruWEp7EwGUshdLKVRCxZW12fp9r40E1mQrf/34g==", + "license": "BSD-3-Clause" + }, + "node_modules/@protobufjs/eventemitter": { + "version": "1.1.1", + "resolved": "https://registry.npmjs.org/@protobufjs/eventemitter/-/eventemitter-1.1.1.tgz", + "integrity": "sha512-vW1GmwMZNnL+gMRaovlh9yZX74kc+TTU3FObkkurpMaRtBfLP3ldjS9KQWlwZgraRE0+dheEEoAxdzcJQ8eXZg==", + "license": "BSD-3-Clause" + }, + "node_modules/@protobufjs/fetch": { + "version": "1.1.1", + "resolved": "https://registry.npmjs.org/@protobufjs/fetch/-/fetch-1.1.1.tgz", + "integrity": "sha512-GpptLrs57adMSuHi3VNj0mAF8dwh36LMaYF6XyJ6JMWlVsc+t42tm1HSEDmOs3A8fC9yyeisgLhsTVQokOZ0zw==", + "license": "BSD-3-Clause", + "dependencies": { + "@protobufjs/aspromise": "^1.1.1" + } + }, + "node_modules/@protobufjs/float": { + "version": "1.0.2", + "resolved": "https://registry.npmjs.org/@protobufjs/float/-/float-1.0.2.tgz", + "integrity": "sha512-Ddb+kVXlXst9d+R9PfTIxh1EdNkgoRe5tOX6t01f1lYWOvJnSPDBlG241QLzcyPdoNTsblLUdujGSE4RzrTZGQ==", + "license": "BSD-3-Clause" + }, + "node_modules/@protobufjs/path": { + "version": "1.1.2", + "resolved": "https://registry.npmjs.org/@protobufjs/path/-/path-1.1.2.tgz", + "integrity": "sha512-6JOcJ5Tm08dOHAbdR3GrvP+yUUfkjG5ePsHYczMFLq3ZmMkAD98cDgcT2iA1lJ9NVwFd4tH/iSSoe44YWkltEA==", + "license": "BSD-3-Clause" + }, + "node_modules/@protobufjs/pool": { + "version": "1.1.0", + "resolved": "https://registry.npmjs.org/@protobufjs/pool/-/pool-1.1.0.tgz", + "integrity": "sha512-0kELaGSIDBKvcgS4zkjz1PeddatrjYcmMWOlAuAPwAeccUrPHdUqo/J6LiymHHEiJT5NrF1UVwxY14f+fy4WQw==", + "license": "BSD-3-Clause" + }, + "node_modules/@protobufjs/utf8": { + "version": "1.1.2", + "resolved": "https://registry.npmjs.org/@protobufjs/utf8/-/utf8-1.1.2.tgz", + "integrity": "sha512-b1UQwcEZ4yCnMCD8DAL1VlbvBJE9/IX4FTIp7BG1xYpf29SLazLSrqUkj4w7Y5y7cCVP6E5tcqqcI0xemPkHug==", + "license": "BSD-3-Clause" + }, "node_modules/@rollup/plugin-commonjs": { "version": "29.0.0", "resolved": "https://registry.npmjs.org/@rollup/plugin-commonjs/-/plugin-commonjs-29.0.0.tgz", @@ -6211,9 +6269,9 @@ "license": "MIT" }, "node_modules/d3-request/node_modules/d3-dispatch": { - "version": "1.0.3", - "resolved": "https://registry.npmjs.org/d3-dispatch/-/d3-dispatch-1.0.3.tgz", - "integrity": "sha512-Qh2DR3neW3lq/ug4oymXHYoIsA91nYt47ERb+fPKjRg6zLij06aP7KqHHl2NyziK9ASxrR3GLkHCtZvXe/jMVg==", + "version": "1.0.6", + "resolved": "https://registry.npmjs.org/d3-dispatch/-/d3-dispatch-1.0.6.tgz", + "integrity": "sha512-fVjoElzjhCEy+Hbn8KygnmMS7Or0a9sI2UzGwoB7cCtvI1XpVN9GpoYlnb3xt2YV66oXYb1fLJ8GMvP4hdU1RA==", "license": "BSD-3-Clause" }, "node_modules/d3-request/node_modules/d3-dsv": { @@ -7414,6 +7472,12 @@ "node": ">=16" } }, + "node_modules/flatbuffers": { + "version": "25.9.23", + "resolved": "https://registry.npmjs.org/flatbuffers/-/flatbuffers-25.9.23.tgz", + "integrity": "sha512-MI1qs7Lo4Syw0EOzUl0xjs2lsoeqFku44KpngfIduHBYvzm8h2+7K8YMQh1JtVVVrUvhLpNwqVi4DERegUJhPQ==", + "license": "Apache-2.0" + }, "node_modules/flatted": { "version": "3.3.2", "resolved": "https://registry.npmjs.org/flatted/-/flatted-3.3.2.tgz", @@ -7718,6 +7782,12 @@ "dev": true, "license": "MIT" }, + "node_modules/guid-typescript": { + "version": "1.0.9", + "resolved": "https://registry.npmjs.org/guid-typescript/-/guid-typescript-1.0.9.tgz", + "integrity": "sha512-Y8T4vYhEfwJOTbouREvG+3XDsjr8E3kIr7uf+JZ0BYloFsttiHU0WfvANVsR7TxNUJa/WpCnw/Ino/p+DeBhBQ==", + "license": "ISC" + }, "node_modules/has-flag": { "version": "4.0.0", "resolved": "https://registry.npmjs.org/has-flag/-/has-flag-4.0.0.tgz", @@ -7877,9 +7947,9 @@ } }, "node_modules/hyphy-scope": { - "version": "1.9.1", - "resolved": "https://registry.npmjs.org/hyphy-scope/-/hyphy-scope-1.9.1.tgz", - "integrity": "sha512-KzvOkZGhAgnJINThk/PcfAq6xYOVX2RYFeZ9WDuxiNB0r2EXAHOGwHH/LPcHXRS/93R3b+cwTGQuUtEYAId/Pg==", + "version": "1.10.0", + "resolved": "https://registry.npmjs.org/hyphy-scope/-/hyphy-scope-1.10.0.tgz", + "integrity": "sha512-vywKNq2wgAwDns16orJn95tZqiKNLVKzuNywIHglKrebHwec0QENRbNDtimx9J+ux6YpbUOJTmwkMtaa5jZdUQ==", "license": "MIT", "dependencies": { "@observablehq/plot": "^0.6.11", @@ -8652,9 +8722,9 @@ "license": "MIT" }, "node_modules/lodash-es": { - "version": "4.17.21", - "resolved": "https://registry.npmjs.org/lodash-es/-/lodash-es-4.17.21.tgz", - "integrity": "sha512-mKnC+QJ9pWVzv+C4/U3rRsHapFfHvQFoFB92e52xeyGMcX6/OlIl78je1u8vePzYZSkkogMPJ2yjxxsb89cxyw==", + "version": "4.18.1", + "resolved": "https://registry.npmjs.org/lodash-es/-/lodash-es-4.18.1.tgz", + "integrity": "sha512-J8xewKD/Gk22OZbhpOVSwcs60zhd95ESDwezOFuA3/099925PdHJ7OFHNTGtajL3AlZkykD32HykiMo+BIBI8A==", "license": "MIT" }, "node_modules/lodash.castarray": { @@ -8729,6 +8799,12 @@ "node": ">= 12.0.0" } }, + "node_modules/long": { + "version": "5.3.2", + "resolved": "https://registry.npmjs.org/long/-/long-5.3.2.tgz", + "integrity": "sha512-mNAgZ1GmyNhD7AuqnTG3/VQ26o760+ZYBPKjPvugO8+nLbYfX6TVpJPseBvopbdY+qpZ/lKUnmEc1LeZYS3QAA==", + "license": "Apache-2.0" + }, "node_modules/loupe": { "version": "3.2.1", "resolved": "https://registry.npmjs.org/loupe/-/loupe-3.2.1.tgz", @@ -9428,6 +9504,26 @@ "url": "https://github.com/sponsors/sindresorhus" } }, + "node_modules/onnxruntime-common": { + "version": "1.27.0", + "resolved": "https://registry.npmjs.org/onnxruntime-common/-/onnxruntime-common-1.27.0.tgz", + "integrity": "sha512-3KxL5wIVqa8Ex08jxSzncm9CMgw8CjOFyOQ7SxvG9o0cVLlhTNKXyIQuTbtX4tGPJEf73OER2xrjt4HJSBL4ow==", + "license": "MIT" + }, + "node_modules/onnxruntime-web": { + "version": "1.27.0", + "resolved": "https://registry.npmjs.org/onnxruntime-web/-/onnxruntime-web-1.27.0.tgz", + "integrity": "sha512-ogDLsqIozHZwifPuN37OproAo0byX6t43/bP8GzeZWBWD6MOGExswFAx3up4NS/vvWBOg2u2PXomDt3rMmdQSg==", + "license": "MIT", + "dependencies": { + "flatbuffers": "^25.1.24", + "guid-typescript": "^1.0.9", + "long": "^5.2.3", + "onnxruntime-common": "1.27.0", + "platform": "^1.3.6", + "protobufjs": "^7.2.4" + } + }, "node_modules/open": { "version": "8.4.2", "resolved": "https://registry.npmjs.org/open/-/open-8.4.2.tgz", @@ -9851,6 +9947,12 @@ "pathe": "^2.0.1" } }, + "node_modules/platform": { + "version": "1.3.6", + "resolved": "https://registry.npmjs.org/platform/-/platform-1.3.6.tgz", + "integrity": "sha512-fnWVljUchTro6RiCFvCXBbNhJc2NijN7oIQxbwsyL0buWJPG85v81ehlHI9fXrJsMNgTofEoWIQeClKpgxFLrg==", + "license": "MIT" + }, "node_modules/playwright": { "version": "1.57.0", "resolved": "https://registry.npmjs.org/playwright/-/playwright-1.57.0.tgz", @@ -10329,6 +10431,29 @@ "node": ">=0.4.0" } }, + "node_modules/protobufjs": { + "version": "7.6.5", + "resolved": "https://registry.npmjs.org/protobufjs/-/protobufjs-7.6.5.tgz", + "integrity": "sha512-/FPD0nUc9jH6rfFjji9IBqOz4pcSE3CsT1m7Ep6Mdb0LxSUMj8hgl6GomOvZzpNpAqqGaXA0P3VSrZLFzIhQrw==", + "hasInstallScript": true, + "license": "BSD-3-Clause", + "dependencies": { + "@protobufjs/aspromise": "^1.1.2", + "@protobufjs/base64": "^1.1.2", + "@protobufjs/codegen": "^2.0.5", + "@protobufjs/eventemitter": "^1.1.1", + "@protobufjs/fetch": "^1.1.1", + "@protobufjs/float": "^1.0.2", + "@protobufjs/path": "^1.1.2", + "@protobufjs/pool": "^1.1.0", + "@protobufjs/utf8": "^1.1.1", + "@types/node": ">=13.7.0", + "long": "^5.3.2" + }, + "engines": { + "node": ">=12.0.0" + } + }, "node_modules/proxy-from-env": { "version": "1.1.0", "resolved": "https://registry.npmjs.org/proxy-from-env/-/proxy-from-env-1.1.0.tgz", diff --git a/package.json b/package.json index 544a80d..8f9282c 100644 --- a/package.json +++ b/package.json @@ -9,8 +9,8 @@ }, "scripts": { "dev": "vite dev", - "prebuild": "svelte-kit sync", - "build": "svelte-kit sync && vite build", + "prebuild": "node scripts/copy-ort-wasm.mjs && svelte-kit sync", + "build": "node scripts/copy-ort-wasm.mjs && svelte-kit sync && vite build", "preview": "vite preview", "prepare": "svelte-kit sync || echo ''", "check": "svelte-kit sync && svelte-check --tsconfig ./tsconfig.json", @@ -37,6 +37,8 @@ "test:e2e": "npm run test:e2e:fast && npm run test:e2e:slow", "test:e2e:fast": "playwright test --grep-invert @slow", "test:e2e:slow": "playwright test --grep @slow --workers=1", + "verify:axomeme-reference": "python3 scripts/axomeme/verify_preprocessing.py", + "verify:axomeme-preprocessing": "node scripts/axomeme/verify_preprocessing.mjs", "verify:model-parity": "python3 scripts/prescreen/verify_parity.py", "verify:model-parity:self-test": "python3 scripts/prescreen/verify_parity.py --self-test", "verify:model-parity:boundary": "python3 scripts/prescreen/verify_parity.py --boundary", @@ -47,7 +49,9 @@ "deploy:cloudflare": "npx wrangler pages deploy .svelte-kit/cloudflare", "deploy:cloudflare:storybook": "npm run build-storybook && npx wrangler pages deploy storybook-static --project-name datamonkey3-storybook", "deploy:cloudflare:all": "./scripts/deploy-with-storybook.sh", - "dev:cloudflare": "npx wrangler pages dev .svelte-kit/cloudflare" + "dev:cloudflare": "npx wrangler pages dev .svelte-kit/cloudflare", + "ort:wasm": "node scripts/copy-ort-wasm.mjs", + "predev": "node scripts/copy-ort-wasm.mjs" }, "devDependencies": { "@chromatic-com/storybook": "^4.1.1", @@ -101,10 +105,11 @@ "@sentry/sveltekit": "^9", "alivibe": ">=0.1.0", "d3": "^7.9.0", - "hyphy-scope": "^1.9.1", + "hyphy-scope": "^1.10.0", "jszip": "^3.10.1", "lucide-svelte": "^0.559.0", "marked": "^15.0.12", + "onnxruntime-web": "^1.27.0", "phylotree": "^2.2.1", "socket.io-client": "^4.8.1", "socket.io-stream": "^0.9.1", diff --git a/scripts/axomeme/compare_to_reference.mjs b/scripts/axomeme/compare_to_reference.mjs new file mode 100644 index 0000000..686d7d2 --- /dev/null +++ b/scripts/axomeme/compare_to_reference.mjs @@ -0,0 +1,192 @@ +/** + * Compare the JS AxoMEME pipeline against the reference implementation, per site. + * + * WHAT THIS DOES AND DOES NOT MEASURE. This is a PORT-CORRECTNESS check, not a model-quality one. + * It runs both implementations over the same alignments and compares their per-site output, so a + * disagreement means our JS differs from their Python. It says nothing about whether the model is + * any good — for that you need MEME's own results on data the model never trained on, and most of + * our corpus is AxoMEME fine-tuning data. + * + * The reference CSVs must come from a TOKENIZER-PATCHED copy of predict_regression_nexus.py. Its + * shipped tokenizer disagrees with the model's training vocabulary on 63 of 64 codons, so comparing + * against it unpatched would measure that bug rather than this port. + * + * Usage: + * node scripts/axomeme/compare_to_reference.mjs --refdir --data + */ + +import { readFileSync, readdirSync, existsSync } from 'node:fs'; +import { join } from 'node:path'; +import * as ort from 'onnxruntime-web/wasm'; +import { parseAlignment } from '../../src/lib/utils/fastaValidation.js'; +import { prepareAlignment, batchSizeFor } from '../../src/lib/services/axomeme/assemble.js'; +import { buildPredictions, siteVariability } from '../../src/lib/services/axomeme/postprocess.js'; + +const arg = (f, d) => { + const i = process.argv.indexOf(f); + return i > 0 ? process.argv[i + 1] : d; +}; +const refDir = arg('--refdir'); +const dataDir = arg('--data'); +const fastaDir = arg('--fasta'); +const modelPath = arg('--model'); +if (!refDir || !dataDir || !modelPath || !fastaDir) { + console.error('usage: --refdir --data --fasta --model '); + process.exit(2); +} + +ort.env.wasm.numThreads = 1; +const session = await ort.InferenceSession.create(new Uint8Array(readFileSync(modelPath))); + +/** Pearson correlation. */ +function pearson(a, b) { + const n = a.length; + if (n < 2) return NaN; + const ma = a.reduce((x, y) => x + y, 0) / n; + const mb = b.reduce((x, y) => x + y, 0) / n; + let num = 0, + da = 0, + db = 0; + for (let i = 0; i < n; i++) { + const x = a[i] - ma, + y = b[i] - mb; + num += x * y; + da += x * x; + db += y * y; + } + return da > 0 && db > 0 ? num / Math.sqrt(da * db) : NaN; +} + +/** Spearman: Pearson on average-tied ranks — the metric the ML team reports. */ +function spearman(a, b) { + const rank = (v) => { + const idx = v.map((x, i) => i).sort((i, j) => v[i] - v[j]); + const r = new Array(v.length); + let i = 0; + while (i < v.length) { + let j = i; + while (j + 1 < v.length && v[idx[j + 1]] === v[idx[i]]) j++; + const avg = (i + j) / 2 + 1; + for (let k = i; k <= j; k++) r[idx[k]] = avg; + i = j + 1; + } + return r; + }; + return pearson(rank(a), rank(b)); +} + +const rows = []; +for (const f of readdirSync(refDir) + .filter((x) => x.endsWith('.csv')) + .sort()) { + const id = f.replace(/\.csv$/, ''); + // FASTA dumped by the REFERENCE's own parser, so both sides see byte-identical sequences and a + // disagreement is a pipeline disagreement rather than two parsers reading a file differently. + const fa = join(fastaDir, `${id}.fa`); + const tre = join(dataDir, `${id}.tre`); + if (!existsSync(fa) || !existsSync(tre)) continue; + + // --- reference --- + const csv = readFileSync(join(refDir, f), 'utf8').trim().split(/\r?\n/); + const header = csv[0].split(','); + const cLrt = header.indexOf('predicted_lrt'); + const cVar = header.indexOf('is_variable'); + if (cLrt < 0) { + console.log(`${id}: no predicted_lrt column`); + continue; + } + const refLrt = [], + refVar = []; + for (let i = 1; i < csv.length; i++) { + const p = csv[i].split(','); + refLrt.push(Number(p[cLrt])); + refVar.push(Number(p[cVar])); + } + + // --- ours --- + let mine; + try { + const parsed = parseAlignment(readFileSync(fa, 'utf8')); + const names = parsed.sequences.map((s) => s.header); + const seqs = parsed.sequences.map((s) => s.sequence); + const prep = prepareAlignment({ names, sequences: seqs, treeText: readFileSync(tre, 'utf8') }); + const size = batchSizeFor(prep.speciesCount); + const acc = { lrt: [], alpha: [], beta_neg: [], beta_pos: [], p_neg: [] }; + for (let s = 0; s < prep.totalCodons; s += size) { + const b = prep.batch(s, size); + const out = await session.run({ + msa_codons: new ort.Tensor('int64', b.msa_codons.data, b.msa_codons.dims), + msa_aas: new ort.Tensor('int64', b.msa_aas.data, b.msa_aas.dims), + dist_matrix: new ort.Tensor('float32', b.dist_matrix.data, b.dist_matrix.dims), + mds_coords: new ort.Tensor('float32', b.mds_coords.data, b.mds_coords.dims), + padding_mask: new ort.Tensor('bool', b.padding_mask.data, b.padding_mask.dims) + }); + for (const k of Object.keys(acc)) acc[k].push(...out[k].data); + } + const selected = prep.selectedNames.map((n) => seqs[names.indexOf(n)] ?? ''); + const variable = siteVariability(selected, prep.totalCodons); + const refSeq = seqs[names.indexOf(prep.referenceName)] ?? seqs[0]; + const refCodons = Array.from({ length: prep.totalCodons }, (_, i) => + refSeq.slice(i * 3, i * 3 + 3) + ); + mine = buildPredictions(acc, { refCodons, variable }); + } catch (e) { + console.log(`${id}: JS threw — ${e.message}`); + continue; + } + + const n = Math.min(mine.length, refLrt.length); + if (n === 0 || mine.length !== refLrt.length) { + console.log(`${id}: SITE COUNT MISMATCH js=${mine.length} ref=${refLrt.length}`); + continue; + } + // Compare only sites BOTH sides consider variable; invariant sites are hard zeros on both and + // would inflate every correlation toward 1. + const ja = [], + ra = []; + let worst = 0, + varMismatch = 0; + for (let i = 0; i < n; i++) { + if (Boolean(mine[i].isVariable) !== Boolean(refVar[i])) varMismatch++; + if (!mine[i].isVariable || !refVar[i]) continue; + ja.push(mine[i].lrt); + ra.push(refLrt[i]); + worst = Math.max(worst, Math.abs(mine[i].lrt - refLrt[i])); + } + rows.push({ + id, + sites: n, + compared: ja.length, + varMismatch, + worst, + pearson: pearson(ja, ra), + spearman: spearman(ja, ra) + }); + console.log( + `${id} sites=${String(n).padStart(5)} compared=${String(ja.length).padStart(5)}` + + ` varMismatch=${String(varMismatch).padStart(4)} max|Δ|=${worst.toExponential(2)}` + + ` r=${pearson(ja, ra).toFixed(6)} rho=${spearman(ja, ra).toFixed(6)}` + ); +} + +console.log('\n=== SUMMARY ==='); +console.log(`alignments compared: ${rows.length}`); +if (rows.length) { + const meanBy = (k) => + rows.reduce((s, r) => s + (Number.isFinite(r[k]) ? r[k] : 0), 0) / rows.length; + console.log(`total sites: ${rows.reduce((s, r) => s + r.sites, 0).toLocaleString()}`); + console.log( + `variable-site classification mismatches: ${rows.reduce((s, r) => s + r.varMismatch, 0)}` + ); + console.log( + `worst per-site |Δ| overall: ${Math.max(...rows.map((r) => r.worst)).toExponential(3)}` + ); + console.log(`mean Pearson r: ${meanBy('pearson').toFixed(6)}`); + console.log(`mean Spearman ρ: ${meanBy('spearman').toFixed(6)}`); + const bad = rows.filter((r) => !(r.pearson > 0.99)); + if (bad.length) { + console.log(`\n[!] ${bad.length} alignments with r <= 0.99:`); + for (const b of bad) + console.log(` ${b.id} r=${b.pearson.toFixed(4)} max|Δ|=${b.worst.toExponential(2)}`); + } +} diff --git a/scripts/axomeme/verify_preprocessing.mjs b/scripts/axomeme/verify_preprocessing.mjs new file mode 100644 index 0000000..cbab363 --- /dev/null +++ b/scripts/axomeme/verify_preprocessing.mjs @@ -0,0 +1,193 @@ +/** + * Compare the JS patristic-distance port against the ML team's own Python, over real trees. + * + * This is the gate that makes the port trustworthy. The unit tests in + * src/test/axomeme-patristic.test.js check hand-computed values on trees small enough to verify by + * eye; this checks the SAME CODE against the reference implementation on real DataMonkey + * submissions, where the trees are large, unbalanced, occasionally malformed, and carry the negative + * branch lengths DM3's own NJ produces. + * + * Run verify_preprocessing.py first to produce the reference JSON. + * + * node scripts/axomeme/verify_preprocessing.mjs \ + * --reference /tmp/axomeme_reference.json \ + * [--tolerance 1e-9] + * + * Exits non-zero if any tree exceeds the tolerance, so it can be wired into CI the way + * scripts/prescreen/verify_parity.py is. Prints file ids, shapes and deltas only — never taxon + * names or sequence data, because the corpus is unpublished research data. + */ + +import { readFileSync } from 'node:fs'; +import { + parseNewick, + leafIndex, + normalizeTaxonName +} from '../../src/lib/services/axomeme/newick.js'; +import { patristicMatrix } from '../../src/lib/services/axomeme/patristic.js'; +import { computeMdsCoordinates } from '../../src/lib/services/axomeme/mds.js'; + +const arg = (flag, fallback) => { + const i = process.argv.indexOf(flag); + return i > 0 && process.argv[i + 1] ? process.argv[i + 1] : fallback; +}; + +const referencePath = arg('--reference'); +const tolerance = Number(arg('--tolerance', '1e-9')); +if (!referencePath) { + console.error('usage: verify_preprocessing.mjs --reference [--tolerance 1e-9]'); + process.exit(2); +} + +const reference = JSON.parse(readFileSync(referencePath, 'utf8')); +const entries = Object.entries(reference); +console.log(`[*] ${entries.length} reference matrices, tolerance ${tolerance}`); + +let checked = 0; +let cells = 0; +let worst = 0; +let worstFile = null; +const failures = []; +const skipped = []; + +/** + * MDS gets its own, much looser tolerance, and the reason is not slack. + * + * The coordinates are cast to float32 by the reference, so ~1e-7 relative is the floor no matter how + * carefully either side computes. On top of that, this is the one stage where two CORRECT + * implementations can legitimately disagree: numpy uses divide-and-conquer (dsyevd), this port uses + * implicit-shift QL, and inside a degenerate eigenspace any orthonormal basis is a valid answer that + * the reference's sign convention cannot disambiguate. So the number below is a measurement + * threshold, not a correctness proof — what matters is the DISTRIBUTION printed at the end. + */ +const mdsTolerance = Number(arg('--mds-tolerance', '1e-5')); +let mdsChecked = 0; +let mdsCells = 0; +let mdsWorstOverall = 0; +let mdsWorstFile = null; +const mdsFailures = []; + +for (const [path, ref] of entries) { + const id = path.split('/').pop(); + let tree; + try { + tree = parseNewick(readFileSync(path, 'utf8')); + } catch (e) { + skipped.push(`${id}: parse threw ${e.message}`); + continue; + } + + const { index } = leafIndex(tree); + // Order the rows exactly as the reference did, so cell (i,j) means the same pair on both sides. + const nodes = ref.names.map((n) => index.get(normalizeTaxonName(n))); + if (nodes.some((v) => v === undefined)) { + // A name the JS parser did not produce is a PARSE divergence, not a distance divergence, and + // it is a real failure — it means the two parsers disagree about what the taxa are. + const missing = nodes.filter((v) => v === undefined).length; + failures.push(`${id}: ${missing}/${ref.names.length} taxa missing from the JS parse`); + continue; + } + + const n = nodes.length; + const mine = patristicMatrix(tree, nodes); + let fileWorst = 0; + for (let i = 0; i < n; i++) { + for (let j = 0; j < n; j++) { + const d = Math.abs(mine[i * n + j] - ref.dist[i][j]); + if (d > fileWorst) fileWorst = d; + } + } + cells += n * n; + checked++; + if (fileWorst > worst) { + worst = fileWorst; + worstFile = `${id} (${n} taxa)`; + } + if (fileWorst > tolerance) { + failures.push(`${id}: ${n} taxa, max|Δ| ${fileWorst.toExponential(3)}`); + } + + // --- MDS, the piece whose agreement is not guaranteed by construction --- + if (!ref.mds) continue; + const cap = ref.max_species; + // The reference pads to max_species BEFORE the eigendecomposition, so the padded zeros take part + // in the double-centring. Build the same padded matrix rather than running MDS on the real taxa. + const padded = new Float64Array(cap * cap); + for (let i = 0; i < n; i++) { + for (let j = 0; j < n; j++) padded[i * cap + j] = mine[i * n + j]; + } + const coords = computeMdsCoordinates(padded, cap, 4); + let mdsWorst = 0; + let mdsWorstComponent = -1; + let mdsScale = 0; + for (let i = 0; i < cap; i++) { + for (let c = 0; c < 4; c++) { + const d = Math.abs(coords[i * 4 + c] - ref.mds[i][c]); + if (d > mdsWorst) { + mdsWorst = d; + mdsWorstComponent = c; + } + const s = Math.abs(ref.mds[i][c]); + if (s > mdsScale) mdsScale = s; + } + } + // GATE ON RELATIVE ERROR, against the largest coordinate in the whole matrix. + // + // An absolute threshold is the wrong instrument here and reads as a failure when nothing is wrong. + // The four components have wildly different scales — a measured 135-taxon tree runs 9.3e2, 2.2e2, + // 6.4e0, 2.1e-1 — because they carry eigenvalues seven orders of magnitude apart. A fixed absolute + // tolerance is simultaneously far too loose for component 0 and far too tight for component 3. + // What the model actually consumes is the vector as a whole, through one Linear layer, so error + // relative to the vector's own scale is the quantity that means something. + const mdsRel = mdsScale > 0 ? mdsWorst / mdsScale : 0; + mdsChecked++; + mdsCells += cap * 4; + if (mdsRel > mdsWorstOverall) { + mdsWorstOverall = mdsRel; + mdsWorstFile = `${id} (${n} taxa, component ${mdsWorstComponent}, abs ${mdsWorst.toExponential(2)})`; + } + if (mdsRel > mdsTolerance) { + mdsFailures.push( + `${id}: ${n} taxa -> ${cap}, rel ${mdsRel.toExponential(3)} ` + + `(abs ${mdsWorst.toExponential(3)}, scale ${mdsScale.toExponential(3)}) on component ${mdsWorstComponent}` + ); + } +} + +console.log(`\n--- patristic distances ---`); +console.log(`[*] compared ${checked} trees, ${cells.toLocaleString()} distance cells`); +console.log(`[*] worst |Δ| ${worst.toExponential(3)}${worstFile ? ` (${worstFile})` : ''}`); +if (skipped.length) { + console.log(`[!] ${skipped.length} skipped:`); + for (const s of skipped.slice(0, 5)) console.log(` ${s}`); +} + +if (mdsChecked) { + console.log(`\n--- MDS coordinates ---`); + console.log(`[*] compared ${mdsChecked} trees, ${mdsCells.toLocaleString()} coordinates`); + console.log( + `[*] worst RELATIVE Δ ${mdsWorstOverall.toExponential(3)}${mdsWorstFile ? ` (${mdsWorstFile})` : ''}` + ); + console.log( + `[*] float32 relative resolution is ~1.2e-7, so anything near that is the format, not the port` + ); + console.log( + `[*] ${mdsChecked - mdsFailures.length}/${mdsChecked} trees within ${mdsTolerance}` + + ` (${(((mdsChecked - mdsFailures.length) / mdsChecked) * 100).toFixed(1)}%)` + ); + if (mdsFailures.length) { + console.log(`[!] ${mdsFailures.length} over tolerance:`); + for (const f of mdsFailures.slice(0, 20)) console.log(` ${f}`); + } +} + +if (failures.length) { + console.log(`\n[!] ${failures.length} DISTANCE FAILURES:`); + for (const f of failures.slice(0, 20)) console.log(` ${f}`); + process.exit(1); +} +if (mdsFailures.length) { + console.log('\n[!] MDS PARITY NOT CLEAN'); + process.exit(1); +} +console.log('\n[*] PARITY OK'); diff --git a/scripts/axomeme/verify_preprocessing.py b/scripts/axomeme/verify_preprocessing.py new file mode 100644 index 0000000..a805731 --- /dev/null +++ b/scripts/axomeme/verify_preprocessing.py @@ -0,0 +1,125 @@ +#!/usr/bin/env python3 +""" +Emit REFERENCE patristic distances for a set of newick trees, for the JS port to be checked against. + +WHY THIS RUNS THE ML TEAM'S CODE RATHER THAN A TRANSCRIPTION OF IT. + +A parity harness that compares our JS against our Python re-implementation proves only that we made +the same mistake twice. So this does not reimplement anything: it extracts the SOURCE TEXT of +`calculate_patristic_distances` out of the handoff's predict_regression_nexus.py and execs it. What +runs is byte-for-byte their function. + +It is extracted rather than imported because importing that module pulls torch, pandas, scipy and +sklearn at module scope, none of which this comparison needs — and requiring a GPU-class dependency +tree to check a tree walk is how a parity gate stops being run. + +The companion is verify_preprocessing.mjs, which consumes this output. Neither script writes anything +into the repository, and neither prints taxon names or sequence data: real DataMonkey submissions are +unpublished research, so only file ids, shapes and aggregate deltas are reported. + +Usage: + python3 scripts/axomeme/verify_preprocessing.py \ + --handoff /path/to/predict_regression_nexus.py \ + --trees '/path/to/corpus/**/*.tre' \ + --out /tmp/axomeme_reference.json +""" + +import argparse +import glob +import json +import math +import re +import sys + +REFERENCE_FNS = ["calculate_patristic_distances", "compute_mds_coordinates"] + + +def extract_reference(handoff_path, fn_name): + """Pull one reference function's source out of the handoff, verbatim.""" + src = open(handoff_path, "r").read() + start = src.find(f"def {fn_name}(") + if start < 0: + sys.exit(f"[!] {fn_name} not found in {handoff_path}") + # Runs to the next top-level def/class. + rest = src[start:] + m = re.search(r"\n(?=(?:def |class )\w)", rest) + return rest[: m.start()] if m else rest + + +def main(): + ap = argparse.ArgumentParser() + ap.add_argument("--handoff", required=True, help="path to predict_regression_nexus.py") + ap.add_argument("--trees", required=True, help="glob for newick files") + ap.add_argument("--out", required=True) + ap.add_argument("--limit", type=int, default=0) + ap.add_argument( + "--max-species", + type=int, + default=512, + help="pad the distance matrix to this size before MDS, as the driver does (default 512)", + ) + args = ap.parse_args() + + from Bio import Phylo # noqa: E402 + import numpy as np # noqa: E402 + + ns = {"math": math, "np": np} + for fn in REFERENCE_FNS: + body = extract_reference(args.handoff, fn) + exec(compile(body, "", "exec"), ns) + print(f"[*] Extracted {fn} ({body.count(chr(10))} lines) from {args.handoff}") + reference = ns["calculate_patristic_distances"] + reference_mds = ns["compute_mds_coordinates"] + + paths = sorted(glob.glob(args.trees, recursive=True)) + if args.limit: + paths = paths[: args.limit] + print(f"[*] {len(paths)} tree files") + + out, crashed, unreadable = {}, [], [] + for p in paths: + try: + tree = Phylo.read(p, "newick") + except Exception as ex: + unreadable.append((p, type(ex).__name__)) + continue + try: + names, dist, _ = reference(tree) + except Exception as ex: + # EXPECTED for a real slice of the corpus, and worth counting rather than skipping: the + # reference computes log((node_count + 1.0) / (dist + 0.1)) unconditionally, so any + # patristic distance <= -0.1 is a math domain error. DM3's own NJ emits the negative + # branch lengths that produce those. + crashed.append((p, type(ex).__name__, str(ex)[:80])) + continue + names = [n for n in names if n] + # The driver pads the distance matrix to max_species BEFORE running MDS, so the padded zeros + # take part in the double-centring and the coordinates depend on max_species. Reproduce that. + # Trees larger than the cap would go through Max-PD selection in the driver; here they are + # simply truncated, identically on both sides, because Max-PD is covered by unit tests and + # mixing the two would make a failure ambiguous. + names = names[: args.max_species] + n = args.max_species + padded = np.zeros((n, n), dtype=np.float32) + for i, a in enumerate(names): + row = dist.get(a, {}) + for j, b in enumerate(names): + padded[i, j] = row.get(b, 0.0) + coords = reference_mds(padded, n_components=4) + out[p] = { + "names": names, + "dist": [[dist[a].get(b, 0.0) for b in names] for a in names], + "max_species": n, + "mds": [[float(v) for v in row] for row in coords], + } + + json.dump(out, open(args.out, "w")) + print(f"[*] wrote {len(out)} reference matrices to {args.out}") + print(f"[*] unreadable by Bio.Phylo: {len(unreadable)}") + print(f"[*] REFERENCE CRASHED on {len(crashed)} trees ({100.0 * len(crashed) / max(1, len(paths)):.1f}%)") + for p, kind, msg in crashed[:5]: + print(f" {p.split('/')[-1]}: {kind}: {msg}") + + +if __name__ == "__main__": + main() diff --git a/scripts/copy-ort-wasm.mjs b/scripts/copy-ort-wasm.mjs new file mode 100644 index 0000000..df1a3aa --- /dev/null +++ b/scripts/copy-ort-wasm.mjs @@ -0,0 +1,58 @@ +/** + * Copy onnxruntime-web's WASM runtime into static/ so DataMonkey serves it itself. + * + * WHY THIS EXISTS. onnxruntime-web does not bundle its WASM binary; it fetches it at runtime, and + * with no configuration it resolves to a jsDelivr CDN URL. That breaks two things at once: + * + * 1. THE PROJECT'S CORE CONSTRAINT. CLAUDE.md: "The website needs to be entirely self-contained so + * that it can be served locally. No pulling down from other domains." A CDN fetch is exactly + * that, and it fails closed on an air-gapped or offline install rather than degrading. + * 2. It fails in dev anyway, which is how this was found — the runtime aborts with "both async and + * sync fetching of the wasm failed" and onnxruntime reports "no available backend found", + * which reads like a broken model rather than a missing asset. + * + * WHY COPY RATHER THAN COMMIT. The file is 12.9 MB of third-party build output that is already + * pinned by package.json and present in node_modules. Committing it would put a binary in git + * history that goes stale the moment onnxruntime-web is upgraded, and nothing would notice. Copying + * at build time means the served asset always matches the installed package. + * + * ONLY THE BASE VARIANT IS COPIED. onnxruntime-web ships ~128 MB of dist covering four builds — + * jsep (WebGPU), asyncify, jspi and the plain SIMD+threads one. AxoMEME runs on CPU with a single + * thread, so it needs the last of those and nothing else. Copying all of them would multiply the + * deploy size for capabilities this feature does not use. + */ + +import { copyFileSync, mkdirSync, existsSync, statSync } from 'node:fs'; +import { dirname, join } from 'node:path'; +import { fileURLToPath } from 'node:url'; + +const repo = join(dirname(fileURLToPath(import.meta.url)), '..'); +const from = join(repo, 'node_modules', 'onnxruntime-web', 'dist'); +const to = join(repo, 'static', 'ort'); + +/** The plain SIMD+threads build and its loader. Nothing else is used at runtime. */ +const ASSETS = ['ort-wasm-simd-threaded.wasm', 'ort-wasm-simd-threaded.mjs']; + +if (!existsSync(from)) { + // Not fatal: `npm run build` in an environment without the optional dependency should say so + // clearly rather than emit a site whose AxoMEME path fails at the last moment. + console.warn('[ort-wasm] onnxruntime-web is not installed; skipping. AxoMEME will not run.'); + process.exit(0); +} + +mkdirSync(to, { recursive: true }); +let total = 0; +for (const name of ASSETS) { + const src = join(from, name); + if (!existsSync(src)) { + console.error( + `[ort-wasm] expected ${name} in onnxruntime-web/dist — has the package layout changed?` + ); + process.exit(1); + } + copyFileSync(src, join(to, name)); + total += statSync(src).size; +} +console.log( + `[ort-wasm] copied ${ASSETS.length} files (${(total / 1024 / 1024).toFixed(1)} MB) to static/ort/` +); diff --git a/scripts/prescreen/requirements.txt b/scripts/prescreen/requirements.txt index 5977e2d..0108924 100644 --- a/scripts/prescreen/requirements.txt +++ b/scripts/prescreen/requirements.txt @@ -1,10 +1,16 @@ # Verification-only dependencies for the MEME hit-likelihood gate. # -# These are NOT application dependencies and must never become one. DM3 ships no ML runtime at all: +# These are NOT application dependencies and must never become one. THE GATE ships no ML runtime: # the browser walks XGBoost's own save_model() JSON in about 40 lines of plain JS (xgbEnsemble.js # over meme_gate.json), with no converter and no intermediate format between the file the ML team # exports and the bytes the browser parses. # +# (This file used to say "DM3 ships no ML runtime at all". That stopped being true when AxoMEME 2.0 +# arrived — it is a 3.78 MB transformer and needs onnxruntime-web, which no hand-written walker is +# going to replace. The gate is the opposite case and stays runtime-free: three features, 500 trees, +# reachable from every method, so a runtime behind it would be paid for by all fifteen. The +# reachability guard in src/test/meme-hit-likelihood.test.js is what holds that line now.) +# # xgboost is here for exactly one reason: to be an INDEPENDENT SECOND IMPLEMENTATION of that walk, # reading the same bytes, so CI can prove the shipped JS reproduces the library's own scoring # (scripts/prescreen/verify_parity.py). It must stay on this side of the fence. Nothing here reaches diff --git a/src/app.css b/src/app.css index d348aa0..c668244 100644 --- a/src/app.css +++ b/src/app.css @@ -50,3 +50,29 @@ border: revert; border-radius: revert; } + +/* --------------------------------------------------------------------------------------------- + Layout reserve for the MEME hit-likelihood row. + + These live here, in the ALWAYS-loaded stylesheet, rather than in MemeHitLikelihood.css, because + two components need them and one of them needs them BEFORE the other exists. RunOutlook draws a + pending placeholder that must be exactly as tall as the row it is standing in for: the panel sits + directly above the Run button, so if the reserve is wrong the button moves under a pointer already + travelling towards it. MemeHitLikelihood.css ships in a lazily-loaded chunk, so anything defined + there is unavailable at precisely the moment the placeholder is drawn. + + They used to be duplicated as hand-measured pixel constants in RunOutlook, with a comment + instructing the next person to re-measure by hand whenever the copy changed. The comment recorded + that the value was already 5.5px stale. Now the reserve is stated once and the placeholder is + derived from it with calc(). + --------------------------------------------------------------------------------------------- */ +:root { + --hit-body-reserve: 97.5px; + --hit-header-line: 23.5px; +} + +@media (min-width: 1024px) { + :root { + --hit-body-reserve: 80px; + } +} diff --git a/src/lib/AnalysisResultViewer.svelte b/src/lib/AnalysisResultViewer.svelte index f7ed60c..56bd17b 100644 --- a/src/lib/AnalysisResultViewer.svelte +++ b/src/lib/AnalysisResultViewer.svelte @@ -4,6 +4,7 @@ import { persistentFileStore } from '../stores/fileInfo'; import ExportPanel from './ExportPanel.svelte'; import FelVisualization from './FelVisualization.svelte'; + import AxomemeResults from './AxomemeResults.svelte'; import AnalysisProgress from './AnalysisProgress.svelte'; import { FINAL_HYPHY_EYE_URL } from './config/env'; import { @@ -119,6 +120,11 @@ case 'fel': case 'contrast-fel': return HyphyScopeFel; + case 'axomeme': + // Not a hyphy-scope visualiser — AxoMEME emits per-site predictions rather than a HyPhy + // results document — but it renders through the same slot, which takes any component + // accepting `data`. + return AxomemeResults; case 'meme': return MemeVisualization; case 'absrel': @@ -184,7 +190,9 @@

Loading analysis results...

{:else if error} -
+
Debug: aBSREL Data Structure -
{JSON.stringify(
+											
{JSON.stringify(
 													resultData,
 													null,
 													2
@@ -377,13 +386,19 @@
 							{/if}
 
 							{#if !resultData.tested && !resultData.fits}
-								
{JSON.stringify(resultData, null, 2)}
+
{JSON.stringify(
+										resultData,
+										null,
+										2
+									)}
{/if} {#if isMethodSupported(analysis.method)}
-

View results with automatic data sharing:

+

+ View results with automatic data sharing: +

{/if}
diff --git a/src/lib/AxomemeResults.svelte b/src/lib/AxomemeResults.svelte new file mode 100644 index 0000000..46b42af --- /dev/null +++ b/src/lib/AxomemeResults.svelte @@ -0,0 +1,82 @@ + + +{#if data} +
+ {#if hasCaveats} +
+
+ + What the model was given +
+
    + {#if summary.matchedFromTree === false} + +
  • + No sequence name matched a tree label, so the tree contributed + nothing — every pair of sequences was treated as equally related. These rankings come + from the alignment alone. Check that your tree's labels match your sequence names. +
  • + {/if} + {#if summary.speciesUsed !== summary.speciesInAlignment} +
  • + {summary.speciesUsed} of {summary.speciesInAlignment} sequences were used — the rest were + not found in the tree{summary.speciesUsed === 512 + ? ', or fell outside the 512-taxon limit' + : ''}. +
  • + {/if} + {#if summary.duplicateSelections > 0} +
  • + {summary.duplicateSelections} taxon slots repeated the same sequence, which happens when + a tree carries no usable branch lengths. +
  • + {/if} + {#if summary.mostNegativeDistance < -0.001} +
  • + The tree contains negative branch lengths — the most negative pairwise distance was + {summary.mostNegativeDistance.toFixed(4)}. Those distances were treated as zero, which + is what the model was trained on, but a tree this far negative is worth checking. +
  • + {/if} + {#each summary.treeWarnings ?? [] as warning} +
  • {warning}
  • + {/each} +
+
+ {/if} + + +
+{/if} diff --git a/src/lib/MemeHitLikelihood.css b/src/lib/MemeHitLikelihood.css index bb0d356..ea775aa 100644 --- a/src/lib/MemeHitLikelihood.css +++ b/src/lib/MemeHitLikelihood.css @@ -67,13 +67,10 @@ all four together. */ .hit-body { padding-left: 22px; /* Align with text after icon, as in AnalysisTimingEstimate */ - min-height: 97.5px; -} - -@media (min-width: 1024px) { - .hit-body { - min-height: 80px; - } + /* Defined in app.css so RunOutlook's placeholder can derive its height from the same number. + The fallback matters: this file ships in a lazily-loaded chunk and must still be correct if it + is ever rendered somewhere the token is not in scope. */ + min-height: var(--hit-body-reserve, 97.5px); } .hit-detail { diff --git a/src/lib/MemeHitLikelihood.svelte b/src/lib/MemeHitLikelihood.svelte index 903ce23..36f6296 100644 --- a/src/lib/MemeHitLikelihood.svelte +++ b/src/lib/MemeHitLikelihood.svelte @@ -187,12 +187,17 @@ } $: applicable = hasHitLikelihood(method); - $: shown = result && result.status !== STATUS.NOT_APPLICABLE ? result : null; + // `!computing` is load-bearing. `seq` stops a superseded RESPONSE overwriting a newer one, but it + // does not clear what is on screen while the new one is in flight — so selecting a second + // alignment kept rendering the first one's band, and its "about 93%" sentence, against data that + // had nothing to do with it. `computing` was assigned in three places and read in none; this is + // where it earns its keep, by letting the pending branch fire. + $: shown = result && !computing && result.status !== STATUS.NOT_APPLICABLE ? result : null; // MEME is selected but there is nothing to score yet. Split out from `shown` because it is not a // failure and must not be drawn as one — and split out from the pending state because it is not // transient. It is reachable: a rejected upload leaves the method dropdown on screen with no // alignment behind it, and this row then spun "estimating…" for as long as the page was open. - $: idle = result && result.status === STATUS.NOT_APPLICABLE ? result : null; + $: idle = result && !computing && result.status === STATUS.NOT_APPLICABLE ? result : null; $: style = shown && shown.level ? LEVEL[shown.level] : null; $: detailText = shown && shown.status !== STATUS.OK ? joinDetail(shown) : null; $: featureLine = shown && shown.num_seqs !== null ? describeFeatures(shown) : null; @@ -366,8 +371,8 @@

How it behaves the score moves one way only: adding - sequences, adding codons, or deeper branches can raise it and never lower it. That holds - at every input the model can score, not just on average. It is still coarse — it + sequences, adding codons, or deeper branches can raise it and never lower it. That + holds at every input the model can score, not just on average. It is still coarse — it reads {LIVE_FEATURES.length} numbers off your alignment and tree and nothing about the sequences themselves — so read it as a rate over similar alignments rather than a property of yours. diff --git a/src/lib/MethodSelector.svelte b/src/lib/MethodSelector.svelte index 329059e..fd54cc4 100644 --- a/src/lib/MethodSelector.svelte +++ b/src/lib/MethodSelector.svelte @@ -9,10 +9,7 @@ import { treeHasBranchLengths } from './services/prescreen/scope.js'; import { AlertTriangle, Play, Loader2 } from 'lucide-svelte'; import { trackEvent } from './utils/analytics.js'; - import { - countBranchGroups, - contrastFelHasEnoughGroups - } from './utils/branchGroupValidation.js'; + import { countBranchGroups, contrastFelHasEnoughGroups } from './utils/branchGroupValidation.js'; export let methodConfig; export let runMethod = null; @@ -22,6 +19,7 @@ // Supported methods - easy to update when methods are implemented const SUPPORTED_METHODS = [ + 'axomeme', 'b-still', 'fel', 'slac', @@ -40,10 +38,23 @@ // Method info with simplified descriptions and runtime estimates const METHOD_INFO = { + axomeme: { + name: 'AxoMEME', + fullName: 'AxoMEME 2.0 — neural surrogate for MEME', + // The MODEL is still under active development, not the integration. That distinction is + // what the badge is for: a user should know the numbers may move between releases even + // though the code around them is settled. + beta: true, + // Deliberately says PREDICTS. This is a fitted model estimating what MEME would report, + // not MEME, and the dropdown is the first place a user forms that expectation. + shortDescription: 'Predict MEME per-site selection in seconds, in your browser', + supported: true + }, 'b-still': { name: 'B-STILL', fullName: 'Bayesian Significance Test of Invariant Low Likelihoods', - shortDescription: 'Detect invariant sites via posterior probabilities and Empirical Bayes Factors', + shortDescription: + 'Detect invariant sites via posterior probabilities and Empirical Bayes Factors', supported: true }, fel: { @@ -161,6 +172,26 @@ // Method-specific advanced options configurations const METHOD_ADVANCED_OPTIONS = { + // AxoMEME exposes almost nothing on purpose. It is a fitted model with one set of weights: + // there are no branch sets to select (it consumes the whole tree as a distance matrix), no + // rate-variation switch, and no genetic code choice — the code table is baked into the + // model's tokenizer as the universal one. Offering knobs that do not reach the model would be + // worse than offering none. Calling mode is the one real choice, because it is a threshold + // applied AFTER inference and genuinely changes what gets reported. + axomeme: { + callMode: { + type: 'select', + label: 'How to rank sites', + // percentile, not the reference driver's pvalue default. The model's predicted LRT does + // not reach the chi-square gates pvalue compares against — measured across 12 real + // submissions, one site in 662 cleared 3.12 and none cleared 4.45 — so pvalue makes the + // method silent. See CALL_DEFAULTS in postprocess.js. + default: 'percentile', + options: ['percentile', 'zscore', 'pvalue'], + description: + 'percentile and zscore rank sites within this alignment, which is what the model is built to do. pvalue compares against fixed LRT gates (4.45 / 3.12) that its scores rarely reach — it will usually report nothing.' + } + }, fel: { // Branch selection options branchesToTest: { @@ -815,7 +846,8 @@ label: 'Amino Acid Property Set', default: '5PROP', options: ['5PROP', '4PROP', '3PROP', '2PROP', 'Atchley', 'LCAP'], - description: 'Set of amino acid properties to model (5PROP: hydrophobicity, polarity, volume, charge, iso-electric point)' + description: + 'Set of amino acid properties to model (5PROP: hydrophobicity, polarity, volume, charge, iso-electric point)' }, // P-value threshold pValueThreshold: { @@ -960,6 +992,37 @@ // Get current method details $: currentMethod = selectedMethod ? availableMethods.find((m) => m.id === selectedMethod) : null; + /** + * A method that runs only in this tab, with no HyPhy binary and no server job behind it. + * + * Two controls are suppressed for these, and in both cases the reason is the same: the control + * would imply a capability the method does not have. + * - EXECUTION MODE. There is no server-side AxoMEME. Offering "Backend Server" would either be + * silently ignored or would route to a socket with no handler. + * - GENETIC CODE. The model's tokenizer bakes in the universal table (its GENETIC_CODE is a + * hard-coded literal), so a user's choice cannot reach it. It currently defaults to Universal, + * which is right, so the control looks harmless -- but changing it would do nothing and say + * nothing, which is the worst of the three options. + */ + $: browserOnly = Boolean(currentMethod?.config?.browserOnly); + // Keep the reported mode honest for analytics and the run-started toast — but REMEMBER what the + // user chose and give it back. + // + // This used to assign `executionMode = 'local'` with no restore. Picking Backend Server for FEL, + // then browsing to AxoMEME, then returning to FEL left the radio on Local and silently ran the + // next analysis through WASM instead of the server — which for a large dataset is the whole reason + // the server exists. Nothing told the user their choice had been discarded. + let executionModeBeforeBrowserOnly = null; + $: if (browserOnly) { + if (executionMode !== 'local') { + executionModeBeforeBrowserOnly = executionMode; + executionMode = 'local'; + } + } else if (executionModeBeforeBrowserOnly) { + executionMode = executionModeBeforeBrowserOnly; + executionModeBeforeBrowserOnly = null; + } + // Get current method's advanced options $: currentMethodOptions = selectedMethod ? METHOD_ADVANCED_OPTIONS[selectedMethod.toLowerCase()] || {} @@ -978,9 +1041,11 @@ } // Auto-detect data type from uploaded file for GARD - $: if ($fileMetricsStore?.FILE_INFO?.gencodeid !== undefined && + $: if ( + $fileMetricsStore?.FILE_INFO?.gencodeid !== undefined && selectedMethod?.toLowerCase() === 'gard' && - methodOptions[selectedMethod]) { + methodOptions[selectedMethod] + ) { const gencodeid = $fileMetricsStore.FILE_INFO.gencodeid; const detectedType = gencodeid === -2 ? 'protein' : gencodeid === -1 ? 'nucleotide' : 'codon'; if (methodOptions[selectedMethod].datatype !== detectedType) { @@ -1002,30 +1067,30 @@ } // Pre-computed renderable options array (excludes interactive-tree) - $: renderableAdvancedOptions = (selectedMethod && methodOptions[selectedMethod]) - ? Object.entries(currentMethodOptions) - .filter(([_, c]) => c.type !== 'interactive-tree') - .map(([key, config]) => { - const isEnabled = !config.dependsOn || - (methodOptions[selectedMethod] && - config.enabledWhen && - config.enabledWhen.includes( - methodOptions[selectedMethod][config.dependsOn] - )); - - // Resolve filtered options for selects constrained by another option - let effectiveConfig = config; - if (config.filteredOptionsBy && config.filteredOptions) { - const controllerValue = methodOptions[selectedMethod][config.filteredOptionsBy]; - const filtered = config.filteredOptions[controllerValue]; - if (filtered) { - effectiveConfig = { ...config, options: filtered }; - } - } - - return { key, config: effectiveConfig, isEnabled }; - }) - : []; + $: renderableAdvancedOptions = + selectedMethod && methodOptions[selectedMethod] + ? Object.entries(currentMethodOptions) + .filter(([_, c]) => c.type !== 'interactive-tree') + .map(([key, config]) => { + const isEnabled = + !config.dependsOn || + (methodOptions[selectedMethod] && + config.enabledWhen && + config.enabledWhen.includes(methodOptions[selectedMethod][config.dependsOn])); + + // Resolve filtered options for selects constrained by another option + let effectiveConfig = config; + if (config.filteredOptionsBy && config.filteredOptions) { + const controllerValue = methodOptions[selectedMethod][config.filteredOptionsBy]; + const filtered = config.filteredOptions[controllerValue]; + if (filtered) { + effectiveConfig = { ...config, options: filtered }; + } + } + + return { key, config: effectiveConfig, isEnabled }; + }) + : []; // Update genetic code ID when name changes $: { @@ -1047,12 +1112,15 @@ } // RELAX branch validation - requires TEST and REFERENCE branches to be tagged - $: relaxHasTestBranches = selectedMethod?.toLowerCase() === 'relax' && + $: relaxHasTestBranches = + selectedMethod?.toLowerCase() === 'relax' && methodOptions?.relax?.interactiveTree?.includes('{TEST}'); - $: relaxHasReferenceBranches = selectedMethod?.toLowerCase() === 'relax' && + $: relaxHasReferenceBranches = + selectedMethod?.toLowerCase() === 'relax' && (methodOptions?.relax?.interactiveTree?.includes('{REFERENCE}') || - methodOptions?.relax?.referenceBranches === 'All'); - $: relaxBranchesValid = selectedMethod?.toLowerCase() !== 'relax' || + methodOptions?.relax?.referenceBranches === 'All'); + $: relaxBranchesValid = + selectedMethod?.toLowerCase() !== 'relax' || (relaxHasTestBranches && relaxHasReferenceBranches); // Contrast-FEL branch validation - compares two or more groups, so at least two @@ -1089,7 +1157,9 @@ // alongside an inferred NJ tree that does have lengths, and depth is the whole point of the // estimate. The source travels with it so the estimate can disclose that NJ lengths are // nucleotide distances rather than the codon-model lengths MEME fits. - $: outlookTree = pickOutlookTree($treeStore); + // Only for a mounted RunOutlook. pickOutlookTree runs treeHasBranchLengths over the newick on every + // tree-store update, which is wasted work when the panel is suppressed. + $: outlookTree = browserOnly ? { newick: '', source: 'unknown' } : pickOutlookTree($treeStore); function pickOutlookTree(store) { const user = store?.usertree || ''; @@ -1243,6 +1313,7 @@ {#each availableMethods as method} {/each} @@ -1253,6 +1324,14 @@ {#if currentMethod}

{currentMethod.info.shortDescription} + {#if currentMethod.info.beta} +
+ Beta +
+ {/if} {#if !currentMethod.info.supported}
Coming Soon @@ -1262,7 +1341,17 @@ {/if} - {#if selectedMethod && currentMethod?.info.supported} + {#if selectedMethod && currentMethod?.info.supported && browserOnly} +
+

Execution Mode

+

+ Runs in your browser. There is no server-side version of this method, so there is nothing + to choose. The first run downloads about 17 MB of model and runtime; after that scoring + takes a few seconds, and the page is briefly unresponsive while the tree embedding is + computed. +

+
+ {:else if selectedMethod && currentMethod?.info.supported}

Execution Mode

@@ -1304,20 +1393,27 @@ Server temporarily unavailable. Please use Local mode.
{/if} +
+ {/if} - -
- -
+ + {#if selectedMethod && currentMethod?.info.supported && !browserOnly} +
+
{/if} @@ -1330,16 +1426,22 @@ Essential
-
- -
+ {#if browserOnly} +

+ This model reads the standard genetic code. It cannot be changed for this method. +

+ {:else} +
+ +
+ {/if}
@@ -1361,7 +1463,7 @@ updateMethodOption(opt.key, +e.target.value)} + on:input={(e) => updateMethodOption(opt.key, +e.target.value)} min={opt.config.min || 0} max={opt.config.max || 1000} step={opt.config.step || 1} @@ -1374,7 +1476,7 @@ updateMethodOption(opt.key, e.target.checked)} + on:change={(e) => updateMethodOption(opt.key, e.target.checked)} disabled={!opt.isEnabled} /> {opt.config.label} @@ -1384,13 +1486,15 @@ {opt.config.label}: updateMethodOption(opt.key, e.target.value)} + on:input={(e) => updateMethodOption(opt.key, e.target.value)} placeholder={opt.config.placeholder || ''} class="option-input" disabled={!opt.isEnabled} @@ -1526,11 +1630,10 @@ {#if methodOptions?.[selectedMethod]?.branchesToTest === 'Custom'} - Contrast-FEL compares two or more groups — please define at least two branch - sets. + Contrast-FEL compares two or more groups — please define at least two branch sets. {:else} - Contrast-FEL compares two or more groups — please tag at least two branch - groups on the tree. + Contrast-FEL compares two or more groups — please tag at least two branch groups on + the tree. {/if}
@@ -1602,6 +1705,13 @@ margin-bottom: 12px; } + .browser-only-note { + margin: 0; + font-size: 0.8125rem; + line-height: 1.4; + color: #475569; + } + .method-dropdown { flex: 1; padding: 8px 12px; @@ -1635,6 +1745,20 @@ gap: 12px; } + .beta-badge { + display: inline-flex; + align-items: center; + padding: 2px 8px; + background: #ede9fe; + border: 1px solid #8b5cf6; + border-radius: 12px; + font-size: 11px; + font-weight: 500; + color: #5b21b6; + text-transform: uppercase; + letter-spacing: 0.025em; + } + .coming-soon-badge { display: inline-flex; align-items: center; diff --git a/src/lib/RunOutlook.svelte b/src/lib/RunOutlook.svelte index b942a56..f41fa40 100644 --- a/src/lib/RunOutlook.svelte +++ b/src/lib/RunOutlook.svelte @@ -89,20 +89,17 @@ /* Rows are separated by a rule rather than by their own borders, so the panel reads as one object with two lines instead of two objects. */ .outlook-row-pending { - /* Matches the mounted row rather than guessing at it: MemeHitLikelihood's 23.5px header line - plus the body reserve it keeps for its copy, which is 97.5px below the layout's 1024px - breakpoint and 80px above it. MEASURED in the running app — the previous value was 5.5px - short of the real row, which is small but is exactly the kind of nudge that moves the Run - button out from under a pointer already travelling towards it. The panel sits above that - button, so neither the chunk arriving nor the estimate resolving a moment later may change - this height. Change these whenever MemeHitLikelihood's .hit-body reserve changes. */ - min-height: 121px; - } + /* The mounted row's height, DERIVED rather than re-measured: MemeHitLikelihood's header line + plus the body reserve it keeps for its copy. Both are declared once in app.css, which is + always loaded — this placeholder is drawn before the lazily-loaded MemeHitLikelihood chunk + exists, so it cannot read them from that file. - @media (min-width: 1024px) { - .outlook-row-pending { - min-height: 103.5px; - } + Getting this wrong is not cosmetic. The panel sits directly above the Run button, so neither + the chunk arriving nor the estimate resolving a moment later may change this height, or the + button moves out from under a pointer already travelling towards it. The previous version + hard-coded 121px / 103.5px and carried an instruction to re-measure by hand; the media query + is gone because the token already carries the breakpoint. */ + min-height: calc(var(--hit-body-reserve, 97.5px) + var(--hit-header-line, 23.5px)); } .outlook-row + .outlook-row { diff --git a/src/lib/config/methodOptions.toml b/src/lib/config/methodOptions.toml index a22eb5c..dc85e89 100644 --- a/src/lib/config/methodOptions.toml +++ b/src/lib/config/methodOptions.toml @@ -907,3 +907,38 @@ takes_value = true type = "string" default = "" choices = [] + +# AxoMEME is not a HyPhy method and has no CLI. This section exists because the method registry +# reads it, not because there are command-line options to describe: the model is a fixed set of +# weights evaluated in the browser, and the only genuine choice is a threshold applied after +# inference. Notably absent, and absent on purpose: +# - no `code` option. The genetic code is baked into the model's tokenizer as the universal table. +# Offering a choice that cannot reach the model would be worse than offering none. +# - no branch selection. AxoMEME consumes the whole tree as a distance matrix; there is no +# foreground set to tag. +[axomeme] +description = "AxoMEME 2.0 ranks the sites of a codon alignment by how MEME-like their selection signal looks, using a neural model, in seconds instead of hours. It orders sites within an alignment; it does not produce calibrated significance, and it is not a replacement for MEME." + +[[axomeme.options]] +name = "alignment" +description = "An in-frame codon alignment in one of the formats supported by DataMonkey" +takes_value = true +type = "string" +default = "" +choices = [] + +[[axomeme.options]] +name = "tree" +description = "A phylogenetic tree with branch lengths. Branch lengths are required: they become the patristic distance matrix the model reads." +takes_value = true +type = "string" +default = "" +choices = [] + +[[axomeme.options]] +name = "call-mode" +description = "How to rank sites. 'percentile' and 'zscore' order sites within this alignment, which is what the model does well — its authors report Spearman rank correlation, not calibration. 'pvalue' compares the predicted LRT against fixed chi-square gates (4.45 / 3.12) that the model's scores rarely reach, so it will usually report nothing." +takes_value = true +type = "string" +default = "percentile" +choices = ["percentile", "zscore", "pvalue"] diff --git a/src/lib/services/AxomemeAnalysisRunner.js b/src/lib/services/AxomemeAnalysisRunner.js new file mode 100644 index 0000000..2dba619 --- /dev/null +++ b/src/lib/services/AxomemeAnalysisRunner.js @@ -0,0 +1,244 @@ +/** + * AxomemeAnalysisRunner — runs AxoMEME 2.0 entirely in the browser. + * + * WHY THIS IS A THIRD RUNNER RATHER THAN A METHOD IN THE OTHER TWO. The project's checklist for + * adding an analysis is written for HyPhy methods, and every step of it assumes one: a `command` to + * pass to hyphy.wasm, an `outputSuffix` naming the HyPhy JSON it writes, CLI argument mapping in + * WasmAnalysisRunner, a socket event in BackendAnalysisRunner, and a hyphy-eye visualiser keyed on + * the shape of that JSON. AxoMEME has none of those. It is a 3.78 MB neural network evaluated by + * onnxruntime-web against tensors this codebase builds itself, and it emits per-site predictions + * rather than a HyPhy results document. Bending it into WasmAnalysisRunner would mean threading a + * non-HyPhy path through every one of those assumptions. + * + * So it implements the same BaseAnalysisRunner lifecycle — createAnalysis, startAnalysisTracking, + * updateProgress, completeAnalysis — which is the part that actually matters for the UI, and shares + * nothing else. + * + * WHAT IT IS NOT. This is a SURROGATE for MEME, not MEME. It predicts what MEME's per-site + * statistics would be, in seconds instead of hours, and it is wrong sometimes in ways a fitted model + * is not. Nothing in this file should present its output as a completed selection analysis; that + * framing belongs in the UI copy and is why `isSurrogate` is on the result rather than implied. + * + * MAIN-THREAD COST, stated because it is the known weak point: the MDS eigendecomposition runs on the + * padded 512 x 512 matrix regardless of taxon count and takes ~600 ms for a small alignment, ~1.9 s + * at 512 taxa. That is synchronous and blocks rendering. It is tolerable against the hours a real + * MEME run takes, and the loop below yields between phases so progress paints, but a Web Worker is + * the right home for it and is the first thing to do if this feels slow. + */ + +import { BaseAnalysisRunner } from './BaseAnalysisRunner.js'; +import { parseAlignment } from '../utils/fastaValidation.js'; +import { prepareAlignment, batchSizeFor } from './axomeme/assemble.js'; +import { loadSession, runSites, isSessionLoaded } from './axomeme/session.js'; +import { buildPredictions, siteVariability, CALL_DEFAULTS } from './axomeme/postprocess.js'; +import { validateInputBundle, VERIFIED_MODEL_SHA256 } from './axomeme/modelContract.js'; +import { inspectBranchLengths } from '../utils/treeSanitation.js'; + +/** + * Below this magnitude a negative branch length is float noise from NJ's unclamped subtraction, not a + * broken tree. Matches the threshold the results page uses for the distance-clamping notice. + */ +const NOISY_NEGATIVE = -0.001; + +/** Let the browser paint between phases; the heavy stages are synchronous. */ +const yieldToBrowser = () => new Promise((resolve) => setTimeout(resolve, 0)); + +export class AxomemeAnalysisRunner extends BaseAnalysisRunner { + /** + * Analyses the user has cancelled. + * + * The base class cancels by walking `activeAnalyses`, which the other two runners populate with a + * backend job id or a WASM handle. This runner has neither — its work is a plain loop on the main + * thread — so without this a cancelled AxoMEME run kept scoring, kept calling updateProgress + * (re-animating a row the user had cancelled), and finally wrote `completed` over the cancelled + * record and popped a success toast. + */ + cancelledAnalyses = new Set(); + + async cancelAnalysis(analysisId) { + this.cancelledAnalyses.add(analysisId); + await super.cancelAnalysis(analysisId); + } + + /** + * @param {string} method ignored — this runner serves one method + * @param {object} config + * @param {string} fastaData + * @param {string} treeData + * @param {string|null} fileId + */ + async runAnalysis(method, config, fastaData, treeData, fileId = null) { + const analysisId = await this.createAnalysis(fileId, 'AXOMEME'); + // 'wasm', NOT 'browser'. This value is persisted to IndexedDB and read by consumers that know + // only two vocabularies: analysisStore.cleanupInterruptedAnalyses() reaps stale runs by testing + // `executionMode === 'wasm'`, and the result viewer and progress row label the mode from it. A + // third value made an interrupted AxoMEME run unreapable — stuck at `running` forever after a + // refresh during the ~600 ms-2 s synchronous MDS — and rendered "Execution: Unknown". AxoMEME + // runs in the browser like the WASM path, so it shares that label. + this.startAnalysisTracking(analysisId, 'AXOMEME', 'wasm', 'Starting AxoMEME prediction...'); + + try { + this.updateProgress(analysisId, 'parsing', 5, 'Reading alignment...'); + const parsed = parseAlignment(fastaData); + const names = parsed.sequences.map((s) => s.header); + const sequences = parsed.sequences.map((s) => s.sequence); + if (names.length === 0) throw new Error('No sequences found in the alignment'); + await yieldToBrowser(); + + // A tree is not strictly required — the reference falls back to an all-zero distance + // matrix — but the result is close to meaningless, so it is recorded rather than hidden. + const treeReport = treeData ? inspectBranchLengths(treeData) : null; + + this.updateProgress(analysisId, 'preparing', 15, 'Computing tree distances and embedding...'); + await yieldToBrowser(); + const prepared = prepareAlignment({ + names, + sequences, + treeText: treeData || undefined, + maxSpecies: config?.maxSpecies, + referenceName: config?.referenceSequence + }); + if (prepared.totalCodons === 0) { + throw new Error('The reference sequence is shorter than one codon'); + } + + // Loading is 3.78 MB of model plus ~13 MB of runtime on the FIRST run of a session, and + // nothing at all afterwards. Saying which is happening matters — a user who has already run + // one alignment should not be told to wait for a download that will not occur. + const firstLoad = !isSessionLoaded(); + this.updateProgress( + analysisId, + 'loading', + 30, + firstLoad ? 'Downloading the AxoMEME model (first run only)...' : 'Preparing the model...' + ); + const { session, sha256 } = await loadSession(); + // Same CPU-only subpath the session uses. Importing the default entry here instead would + // pull a second, WebGPU-enabled copy of the runtime into the graph. + const ort = await import('onnxruntime-web/wasm'); + + const batch = batchSizeFor(prepared.speciesCount); + // dist_matrix, mds_coords and padding_mask are per-ALIGNMENT — batch() emits b memcpy'd + // copies of one array. Validating every batch rescanned up to the whole 64 MB budget three + // times over, on the main thread, to re-check the same N x N block b times. Validate in full + // once; after that the invariant tensors cannot have changed. + let fullyValidated = false; + const accumulated = { lrt: [], alpha: [], beta_neg: [], beta_pos: [], p_neg: [] }; + for (let start = 0; start < prepared.totalCodons; start += batch) { + const bundle = prepared.batch(start, batch); + // The contract check is cheap next to inference and catches the whole class of errors + // that produce a well-formed tensor meaning something the model never saw. + if (!fullyValidated) { + const check = validateInputBundle(bundle, { + batch: bundle.msa_codons.dims[0], + numSpecies: prepared.speciesCount, + windowSize: prepared.windowSize + }); + if (!check.ok) { + throw new Error(`AxoMEME input check failed: ${check.errors.slice(0, 3).join('; ')}`); + } + fullyValidated = true; + } + const out = await runSites(session, bundle, ort); + // NOT `push(...out[key])`. batchSizeFor scales as 1/N^2, so a few-taxon alignment gets + // batches of 10^5-10^6 sites, and spreading that many elements into a call blows V8's + // argument limit with "Maximum call stack size exceeded" at ~90% progress. Measured: it + // throws at 167,772, which is exactly the batch size for a 10-taxon alignment. + for (const key of Object.keys(accumulated)) { + const src = out[key]; + for (let i = 0; i < src.length; i++) accumulated[key].push(src[i]); + } + + // Between batches is the only place this loop yields, so it is the only place a cancel can + // take effect. Returning without calling completeAnalysis leaves the cancelled record as + // the user left it. + if (this.cancelledAnalyses.has(analysisId)) { + this.cancelledAnalyses.delete(analysisId); + return { analysisId, cancelled: true }; + } + + const done = Math.min(start + batch, prepared.totalCodons); + this.updateProgress( + analysisId, + 'running', + 40 + Math.round((done / prepared.totalCodons) * 50), + `Scoring site ${done} of ${prepared.totalCodons}...` + ); + await yieldToBrowser(); + } + + this.updateProgress(analysisId, 'processing', 95, 'Building per-site results...'); + // Variability is judged over the SELECTED species, which is what the model saw — not over + // every sequence in the file, which may include taxa the tree did not contain. + // By INDEX, not `names.indexOf(name)`. With duplicate FASTA headers indexOf returns the FIRST + // match while orderSpecies keeps the LAST, so the variability flags would come from a + // different sequence than the one the model was tokenised from. + const selectedSeqs = prepared.selectedIndices.map((i) => sequences[i] ?? ''); + const variable = siteVariability(selectedSeqs, prepared.totalCodons); + const refSeq = sequences[prepared.referenceIndex] ?? sequences[0]; + const refCodons = Array.from({ length: prepared.totalCodons }, (_, i) => + refSeq.slice(i * 3, i * 3 + 3) + ); + // ONE source of truth for the calling mode. This and the summary below used to resolve it + // with OPPOSITE precedence, so a config carrying both could score with percentile gates + // while the footer told the reader the calls meant "Z >= 2.5". + const callConfig = { + ...(config?.calling ?? {}), + ...(config?.callMode ? { mode: config.callMode } : {}) + }; + const sites = buildPredictions(accumulated, { refCodons, variable }, callConfig); + + const result = { + method: 'AxoMEME', + modelVersion: '2.0-viral-finetuned', + modelSha256: sha256 ?? VERIFIED_MODEL_SHA256, + // Load-bearing for the UI: these are PREDICTIONS of what MEME would report, not MEME. + isSurrogate: true, + surrogateFor: 'MEME', + sites, + summary: { + totalSites: prepared.totalCodons, + variableSites: sites.filter((s) => s.isVariable).length, + calledSites: sites.filter((s) => s.call !== 'Neutral').length, + speciesUsed: prepared.speciesCount, + speciesInAlignment: names.length, + referenceSequence: prepared.referenceName, + // Named in the footer: it changes what a "call" means, and the default is not the + // reference driver's. + callMode: callConfig.mode ?? CALL_DEFAULTS.mode, + matchedFromTree: prepared.matchedFromTree, + duplicateSelections: prepared.duplicateSelections, + // Reported, not hidden. Clamping matches the model's training pipeline, so it is not + // an error — but a tree whose distances are meaningfully negative is a different + // situation from one carrying float noise, and only the magnitude distinguishes them. + clampedDistances: prepared.clampedDistances, + mostNegativeDistance: prepared.mostNegativeDistance, + // Only surface tree problems worth acting on. inspectBranchLengths reports ANY negative + // branch, and DM3's NJ routinely emits float noise around -1e-5 — the bundled + // large.nex demo has two. Those are clamped (matching the model's training pipeline) + // and warning about them trains users to ignore the warning box, which is worse than + // not having one. A meaningfully negative tree still gets through, via this filter and + // via the mostNegativeDistance line the results page renders. + treeWarnings: + treeReport && !treeReport.ok + ? treeReport.reasons.filter( + (r) => !/negative/.test(r) || (treeReport.min ?? 0) <= NOISY_NEGATIVE + ) + : [] + } + }; + + await this.completeAnalysis(analysisId, true, result); + this.cancelledAnalyses.delete(analysisId); + return { analysisId, result }; + } catch (error) { + await this.completeAnalysis(analysisId, false, null, error.message); + this.cancelledAnalyses.delete(analysisId); + throw error; + } + } +} + +const axomemeAnalysisRunner = new AxomemeAnalysisRunner(); +export default axomemeAnalysisRunner; +export { axomemeAnalysisRunner }; diff --git a/src/lib/services/axomeme/assemble.js b/src/lib/services/axomeme/assemble.js new file mode 100644 index 0000000..800ef5d --- /dev/null +++ b/src/lib/services/axomeme/assemble.js @@ -0,0 +1,322 @@ +/** + * assemble.js — turn an alignment and a tree into the five tensors the AxoMEME graph accepts. + * + * This is the layer that joins everything else: newick parse, patristic distances, Max-PD selection, + * tokenisation and MDS. Two decisions in here are not obvious and both are load-bearing. + * + * --------------------------------------------------------------------------------------------- + * 1. MDS IS COMPUTED ON THE PADDED MATRIX, THEN SLICED. THE MODEL IS FED ONLY THE REAL SPECIES. + * --------------------------------------------------------------------------------------------- + * The reference builds a [max_species, max_species] tensor — 512 x 512 by default — regardless of how + * many taxa the alignment actually has, and runs MDS on that. The padded zeros take part in the + * double-centring, so the coordinates genuinely depend on max_species and computing MDS on the real N + * gives different numbers. That part must be reproduced exactly. + * + * Feeding those 512 slots to the GRAPH, however, is a different question, and the answer is measured + * rather than assumed: feeding N = real species produces the same outputs as feeding N = 512 with the + * remainder masked, to 1.192e-07 — exactly one float32 ulp. That is what masking is for; `forward()` + * excludes padded slots from the validity mask, the distance normalisation, the attention and every + * pooling stream, and softmax's -1e4 fill underflows to zero. + * + * The difference is not academic. A 441-site alignment at N=512 needs + * 441 x 512 x 512 x 4 = 462 MB for dist_matrix alone, because ONNX needs materialised data where + * torch used an `expand` view. At N=36 it is 2.3 MB, and attention is quadratic in N on top of that. + * So: MDS at 512 for fidelity, graph at real N for tractability. + * + * --------------------------------------------------------------------------------------------- + * 2. SPECIES ORDER COMES FROM THE TREE, NOT THE ALIGNMENT, AND THE REFERENCE SEQUENCE LEADS. + * --------------------------------------------------------------------------------------------- + * predict_regression_nexus.py:1069-1082 walks the TREE's leaves in order and keeps the ones present + * in the alignment, then moves the reference sequence to index 0. Order matters twice over: it fixes + * the row order of the distance matrix (and therefore MDS), and index 0 is the seed of the Max-PD + * traversal, so it decides which taxa survive the cap. + * + * The reference sequence itself is picked by a heuristic that is a TOGA-mammal artifact — it looks + * for sequences literally named 'hg', 'hg38' or 'human' before falling back to the first sequence in + * the file. On DataMonkey's viral traffic that fallback is what fires essentially always. + * + * Sources: predict_regression_nexus.py:1049-1082 (reference + matching), 1165-1188 (Max-PD), + * 1213-1275 (tensor allocation, pad values, MDS, the site loop). + */ + +import { parseNewick, leafIndex, normalizeTaxonName } from './newick.js'; +import { patristicMatrix, maxPdSelect } from './patristic.js'; +import { computeMdsCoordinates } from './mds.js'; +import { codonToken, aaToken } from './tokenizer.js'; +import { + MAX_SPECIES_DEFAULT, + WINDOW_SIZE_DEFAULT, + MDS_COMPONENTS, + CODON_UNKNOWN, + AA_UNKNOWN +} from './modelContract.js'; + +/** The reference-sequence heuristic, verbatim from the driver. */ +const REFERENCE_HEURISTICS = ['hg', 'hg38', 'human']; + +/** + * Pick the reference sequence, whose length sets the number of sites and which leads the species + * order. Explicit choice wins; otherwise the driver's heuristic, which falls through to the first + * sequence for anything that is not a TOGA mammal alignment. + * + * @param {string[]} names + * @param {string} [explicit] + * @returns {string} + */ +export function chooseReference(names, explicit) { + if (explicit && names.includes(explicit)) return explicit; + for (const h of REFERENCE_HEURISTICS) if (names.includes(h)) return h; + return names[0]; +} + +/** + * Order the taxa the way the reference does: tree leaf order, filtered to those present in the + * alignment, with the reference sequence moved to the front. + * + * @param {string[]} names alignment sequence names + * @param {any} tree parsed newick, or null + * @param {string} referenceName + * @returns {{order: number[], matchedFromTree: boolean}} indices into `names` + */ +export function orderSpecies(names, tree, referenceName) { + const byNormalised = new Map(); + // Later duplicates overwrite, matching the reference's `name_map[norm] = align_name`. + names.forEach((n, i) => byNormalised.set(normalizeTaxonName(n), i)); + + let order = []; + let matchedFromTree = false; + if (tree) { + const seen = new Set(); + for (const leaf of tree.leaves) { + const nm = normalizeTaxonName(tree.name[leaf]); + if (!nm || seen.has(nm)) continue; + seen.add(nm); + const idx = byNormalised.get(nm); + if (idx !== undefined) order.push(idx); + } + matchedFromTree = order.length > 0; + } + // "Zero matching species between tree and alignment. Falling back to alignment order." + if (order.length === 0) order = names.map((_, i) => i); + + const refIdx = names.indexOf(referenceName); + if (refIdx >= 0) { + const at = order.indexOf(refIdx); + if (at > 0) { + order.splice(at, 1); + order.unshift(refIdx); + } else if (at < 0) { + // The reference is not in the tree. The driver still puts it first, because its sequence + // defines the site count and it seeds Max-PD. + order.unshift(refIdx); + } + } + return { order, matchedFromTree }; +} + +/** + * Prepare everything that is per-alignment rather than per-site. + * + * Returns a handle whose `batch()` materialises tensors for a range of sites. Sites are chunked + * rather than assembled in one go because the tensors are O(sites x N^2): a 12,000-codon upload at + * even 36 species is 62 MB for dist_matrix, and DataMonkey accepts uploads that large. + * + * @param {{names: string[], sequences: string[], tree?: any, treeText?: string, + * maxSpecies?: number, windowSize?: number, referenceName?: string}} input + */ +export function prepareAlignment(input) { + const { + names, + sequences, + treeText, + maxSpecies = MAX_SPECIES_DEFAULT, + windowSize = WINDOW_SIZE_DEFAULT, + referenceName + } = input; + + if (!Array.isArray(names) || !Array.isArray(sequences) || names.length !== sequences.length) { + throw new Error('prepareAlignment: names and sequences must be parallel arrays'); + } + if (names.length === 0) throw new Error('prepareAlignment: no sequences'); + + const tree = input.tree ?? (treeText ? parseNewick(treeText) : null); + const refName = chooseReference(names, referenceName); + const { order, matchedFromTree } = orderSpecies(names, tree, refName); + + // --- taxon selection ------------------------------------------------------------------------- + let selected = order; + let duplicates = 0; + if (order.length > maxSpecies) { + if (tree) { + const { index } = leafIndex(tree); + const nodes = order.map((i) => index.get(normalizeTaxonName(names[i]))); + if (nodes.every((n) => n !== undefined)) { + const pd = maxPdSelect(tree, nodes, maxSpecies); + duplicates = pd.duplicates; + selected = pd.selected.map((k) => order[k]); + } else { + selected = order.slice(0, maxSpecies); + } + } else { + // `selected_species = matching_species[:args.max_species]` + selected = order.slice(0, maxSpecies); + } + } + const n = selected.length; + + // --- distances, at the REAL species count ---------------------------------------------------- + // + // NEGATIVE DISTANCES ARE CLAMPED TO ZERO, and this is the one place in the port that deliberately + // alters a value rather than passing it through. The reasoning matters: + // + // - DM3's own NJ inference emits negative branch lengths. `NJ.bf:214-220` computes the + // three-taxon closed form (d01 + d02 - d12)/2 with no Max(0, ...), so any triangle-inequality + // violation yields one. Most are float noise -- the demo alignment produces a patristic sum of + // -1.04e-5, which is zero with rounding error on it. + // - THE MODEL WAS TRAINED ON CLAMPED DISTANCES. The handoff README is explicit: "This build's + // training pipeline clamps distances >= 0." So clamping here matches what the weights were + // fitted against. It is the reference's INFERENCE path that omits the clamp, and that omission + // is exactly why it throws on 4.6% of real DM3 trees (findings log §2). + // + // So refusing a -1e-5 distance would reject an ordinary tree for a rounding artifact, while + // passing it through unclamped would feed the model something training never showed it. Clamping + // does neither. What is NOT acceptable is doing it silently, so the magnitude is recorded and a + // meaningfully negative tree still reaches the user. + const dist = new Float32Array(n * n); + let clampedDistances = 0; + let mostNegativeDistance = 0; + if (tree) { + const { index } = leafIndex(tree); + const nodes = selected.map((i) => index.get(normalizeTaxonName(names[i]))); + if (nodes.every((v) => v !== undefined)) { + const full = patristicMatrix(tree, nodes); + // float32 because that is what dist_tensor is, and MDS is sensitive to the rounding — see + // mds.js. Doing it here rather than there keeps a single source of the rounded values. + for (let k = 0; k < n * n; k++) { + const v = full[k]; + if (v < 0) { + clampedDistances++; + if (v < mostNegativeDistance) mostNegativeDistance = v; + dist[k] = 0; + } else { + dist[k] = v; + } + } + } + } + // No tree, or names that do not resolve: the reference falls back to an all-zero matrix rather + // than failing, and MDS on all zeros yields all-zero coordinates. + + // --- MDS on the PADDED matrix, then sliced back ---------------------------------------------- + // This is the part that must not be "simplified" to run on `dist` directly. See the header. + const cap = Math.max(maxSpecies, n); + const padded = new Float64Array(cap * cap); + for (let i = 0; i < n; i++) for (let j = 0; j < n; j++) padded[i * cap + j] = dist[i * n + j]; + const paddedMds = computeMdsCoordinates(padded, cap, MDS_COMPONENTS); + const mds = new Float32Array(n * MDS_COMPONENTS); + for (let i = 0; i < n; i++) { + for (let c = 0; c < MDS_COMPONENTS; c++) + mds[i * MDS_COMPONENTS + c] = paddedMds[i * MDS_COMPONENTS + c]; + } + + // --- tokens ---------------------------------------------------------------------------------- + const refSeq = sequences[names.indexOf(refName)] ?? sequences[0]; + const totalCodons = Math.floor(refSeq.length / 3); + const halfWindow = Math.floor(windowSize / 2); + + // Pre-filled with the pad values, so a species whose sequence is shorter than the reference simply + // keeps them — matching `torch.ones(...) * 65` and `* AA_TO_IDX['?']`. + const codonTokens = new BigInt64Array(totalCodons * n * windowSize).fill(BigInt(CODON_UNKNOWN)); + const aaTokens = new BigInt64Array(totalCodons * n * windowSize).fill(BigInt(AA_UNKNOWN)); + + for (let s = 0; s < n; s++) { + const seq = sequences[selected[s]] ?? ''; + const seqCodons = Math.floor(seq.length / 3); + for (let site = 0; site < totalCodons; site++) { + for (let w = 0; w < windowSize; w++) { + // `codon_pos_1based = site_idx - half_win + w_idx`, guarded to the sequence's own length. + const pos = site + 1 - halfWindow + w; + if (pos < 1 || pos > seqCodons) continue; + const codon = seq.slice((pos - 1) * 3, (pos - 1) * 3 + 3); + const at = (site * n + s) * windowSize + w; + codonTokens[at] = BigInt(codonToken(codon)); + aaTokens[at] = BigInt(aaToken(codon)); + } + } + } + + // Every selected slot is a real species, so nothing is padded. The tensor is still required by the + // graph, and is what would carry padding if a caller ever fed padded slots. + const paddingMask = new Uint8Array(n); + + return { + speciesCount: n, + totalCodons, + windowSize, + referenceName: refName, + selectedNames: selected.map((i) => names[i]), + // Indices as well as names. A caller that maps a name back with `names.indexOf()` gets the + // FIRST match, which is the wrong record when the alignment has duplicate headers — and this + // module already resolved duplicates the other way round. + selectedIndices: selected.slice(), + referenceIndex: names.indexOf(refName), + matchedFromTree, + duplicateSelections: duplicates, + clampedDistances, + mostNegativeDistance, + dist, + mds, + paddingMask, + codonTokens, + aaTokens, + + /** + * Materialise the five tensors for sites [start, start + count). + * + * @param {number} start + * @param {number} [count] + * @returns {Record} + */ + batch(start, count = totalCodons - start) { + const b = Math.max(0, Math.min(count, totalCodons - start)); + const perSiteTokens = n * windowSize; + const distData = new Float32Array(b * n * n); + const mdsData = new Float32Array(b * n * MDS_COMPONENTS); + const maskData = new Uint8Array(b * n); + for (let k = 0; k < b; k++) { + distData.set(dist, k * n * n); + mdsData.set(mds, k * n * MDS_COMPONENTS); + maskData.set(paddingMask, k * n); + } + return { + msa_codons: { + data: codonTokens.slice(start * perSiteTokens, (start + b) * perSiteTokens), + dims: [b, n, windowSize] + }, + msa_aas: { + data: aaTokens.slice(start * perSiteTokens, (start + b) * perSiteTokens), + dims: [b, n, windowSize] + }, + dist_matrix: { data: distData, dims: [b, n, n] }, + mds_coords: { data: mdsData, dims: [b, n, MDS_COMPONENTS] }, + padding_mask: { data: maskData, dims: [b, n] } + }; + } + }; +} + +/** + * Site count per batch that keeps the materialised tensors near `budgetBytes`. + * + * dist_matrix dominates at 4 * N^2 bytes per site, so this is a one-term estimate rather than an + * accounting of every tensor. At least one site is always returned — a single site of a 512-taxon + * alignment is 1 MB and has to go through whatever the budget says. + * + * @param {number} speciesCount + * @param {number} [budgetBytes] default 64 MB + * @returns {number} + */ +export function batchSizeFor(speciesCount, budgetBytes = 64 * 1024 * 1024) { + const perSite = 4 * speciesCount * speciesCount; + return Math.max(1, Math.floor(budgetBytes / Math.max(perSite, 1))); +} diff --git a/src/lib/services/axomeme/mds.js b/src/lib/services/axomeme/mds.js new file mode 100644 index 0000000..d920ec2 --- /dev/null +++ b/src/lib/services/axomeme/mds.js @@ -0,0 +1,141 @@ +/** + * mds.js — classical multidimensional scaling, matching predict_regression_nexus.py:161-187. + * + * This produces `mds_coords`, one of the five ONNX graph inputs, and it is the LAST piece of the + * preprocessing port and the only one whose agreement with Python cannot be guaranteed by + * construction. Everything else — newick parse, patristic distances, tokenisation — is exact + * arithmetic over discrete choices. This one runs an eigendecomposition, and eigenvectors are unique + * only up to sign, and inside a degenerate eigenspace not even that. + * + * THE REFERENCE, step for step: + * 1. N <= n_components -> all-zero coordinates, early return. + * 2. cast to float64; D2 = D ** 2 + * 3. H = I - ones(N,N)/N ; B = -0.5 * (H @ D2 @ H) + * 4. evals, evecs = np.linalg.eigh(B) (ascending) + * 5. reorder DESCENDING by eigenvalue + * 6. sign convention: per column, if the largest-magnitude entry is negative, negate the column + * 7. coords[:, i] = evecs[:, i] * sqrt(evals[i]) when evals[i] > 0, else left at zero + * 8. cast to float32 + * + * STEP 6 IS THE GOOD NEWS. The reference already canonicalises eigenvector signs — "largest absolute + * value element is positive". That was going to be an ask to the ML team and turned out to be + * already there, which removes the sign half of the ambiguity for free. What remains is degeneracy: + * when two eigenvalues are equal, any orthonormal basis of their shared eigenspace is a correct + * answer, the sign rule does not disambiguate a rotation, and numpy's divide-and-conquer and our QL + * iteration may land on different bases. That is measured by the parity harness, not argued about + * here. + * + * ONE PLACE THIS DELIBERATELY DOES NOT COPY THE REFERENCE'S ARITHMETIC. The reference forms B with + * two dense matrix multiplications by H. Since H = I - J/N, that product is exactly the + * double-centring identity + * B[i][j] = -0.5 * (D2[i][j] - rowMean[i] - colMean[j] + grandMean) + * so this computes it directly: same value, two O(n^2) passes instead of two O(n^3) matmuls. Bit + * parity with numpy was never available anyway — its matmul is blocked BLAS with FMA, which no + * hand-written loop reproduces — and the output is cast to float32, which is ~1e-7 relative and + * swallows the difference. The harness confirms this rather than assuming it. + * + * MDS RUNS ON THE PADDED MATRIX. The caller must hand this the full [max_species, max_species] + * matrix, zeros included, because the reference does: the padded region participates in the + * double-centring, so coordinates depend on max_species and not only on the real taxa. Running MDS + * on the real N and padding afterwards produces different numbers. See modelContract.js. + */ + +import { symmetricEigen } from './symmetricEigen.js'; +import { MDS_COMPONENTS } from './modelContract.js'; + +/** + * Classical MDS coordinates for a distance matrix. + * + * @param {Float64Array|number[]} dist row-major n*n distance matrix (the PADDED one) + * @param {number} n + * @param {number} [nComponents] + * @returns {Float32Array} n * nComponents, row-major — float32 to match the reference's final cast + */ +export function computeMdsCoordinates(dist, n, nComponents = MDS_COMPONENTS) { + const coords = new Float32Array(n * nComponents); + // `if N <= n_components: return zeros`. Note <=, not <: a 4-taxon alignment with 4 components + // gets all-zero coordinates, which is the reference's behaviour and not an edge case to improve. + if (n <= nComponents) return coords; + + // D2 = D ** 2, WITH THE DISTANCES FIRST ROUNDED TO FLOAT32. + // + // That rounding is not a detail. The reference receives `dist_tensor.numpy()`, and dist_tensor is + // `torch.zeros(max_species, max_species, dtype=torch.float32)` — so the values it squares are + // float32, widened back to float64 by its own `.astype(np.float64)` on the very next line. Passing + // full float64 distances here produces a DIFFERENT MDS, and not subtly: + // + // Real distance matrices are wildly ill-conditioned for this purpose. On a measured 135-taxon + // DM3 tree the eigenvalues run 1.26e7, 1.52e5, 1.30e2, 2.99e-1 — seven orders of magnitude from + // first to fourth. Squared distances reach ~1e6, so a float32 rounding of ~1e-7 relative is an + // ABSOLUTE perturbation of ~0.1 in D2, which is comparable to the fourth eigenvalue itself. + // Components 2 and 3 then differ by 40-99%. Feeding float64 here disagreed with the reference on + // 5 of 270 real trees; rounding to float32 first is what makes them agree. + // + // Squaring also makes negative distances positive, so a negative branch length does not break MDS + // the way it breaks the reference's density term — it silently becomes a positive distance of the + // same magnitude. Worth knowing when reading coordinates from a bad tree. + const D2 = new Float64Array(n * n); + for (let i = 0; i < n * n; i++) { + const v = Math.fround(dist[i]); + D2[i] = v * v; + } + + // Double-centre: B = -0.5 * (H @ D2 @ H), computed via row/column means. D2 is symmetric so the + // column means equal the row means, but they are computed separately anyway — the input is only + // assumed symmetric, and an asymmetric one should degrade predictably rather than silently. + const rowMean = new Float64Array(n); + const colMean = new Float64Array(n); + let grand = 0; + for (let i = 0; i < n; i++) { + let s = 0; + for (let j = 0; j < n; j++) s += D2[i * n + j]; + rowMean[i] = s / n; + grand += s; + } + grand /= n * n; + for (let j = 0; j < n; j++) { + let s = 0; + for (let i = 0; i < n; i++) s += D2[i * n + j]; + colMean[j] = s / n; + } + + const B = new Float64Array(n * n); + for (let i = 0; i < n; i++) { + for (let j = 0; j < n; j++) { + B[i * n + j] = -0.5 * (D2[i * n + j] - rowMean[i] - colMean[j] + grand); + } + } + + const { values, vectors } = symmetricEigen(B, n); + + // eigh returns ascending; the reference reorders descending. Only the top nComponents are used, + // but the SIGN CONVENTION in the reference runs over every column, so it is applied per component + // as they are read rather than to the whole matrix — same result, nComponents passes instead of n. + for (let c = 0; c < nComponents; c++) { + const src = n - 1 - c; // descending order + const value = values[src]; + if (!(value > 0)) continue; // `if val > 0` — non-positive components stay zero + + // Sign convention: the largest-magnitude entry of the column must be positive. np.argmax + // takes the FIRST maximum on a tie, so `>` and not `>=` here. + let maxAbs = -1; + let maxIdx = 0; + for (let i = 0; i < n; i++) { + const a = Math.abs(vectors[i * n + src]); + if (a > maxAbs) { + maxAbs = a; + maxIdx = i; + } + } + // np.sign(0) is 0, and `if sign < 0` is then false — an all-zero column is left alone rather + // than negated. Matching that matters only for degenerate input, but it is free to match. + const flip = vectors[maxIdx * n + src] < 0 ? -1 : 1; + + const scale = Math.sqrt(value); + for (let i = 0; i < n; i++) { + coords[i * nComponents + c] = Math.fround(flip * vectors[i * n + src] * scale); + } + } + + return coords; +} diff --git a/src/lib/services/axomeme/modelContract.js b/src/lib/services/axomeme/modelContract.js new file mode 100644 index 0000000..2e9dbb2 --- /dev/null +++ b/src/lib/services/axomeme/modelContract.js @@ -0,0 +1,335 @@ +/** + * modelContract.js — what the AxoMEME 2.0 ONNX graph accepts and returns. + * + * WHY THIS FILE EXISTS, AND WHY IT IS CODE RATHER THAN A README. + * + * The exported graph is the MODEL ONLY. It takes five already-computed tensors and returns five; + * the entire preprocessing pipeline that produces those tensors — newick parse, patristic + * distances, taxon subsampling, codon/AA tokenisation, and the MDS eigendecomposition — is NOT in + * the graph and has to be rebuilt in JS. `mds_coords` is a graph INPUT, not something the graph + * computes. (`torch.linalg.eigh` has no ONNX lowering, so this is forced rather than chosen: + * exporting the identical module with eigh removed succeeds and with it present fails.) + * + * That means every number below is a place a JS port can be silently, plausibly wrong — producing + * a well-formed tensor of the right shape that means something different from what the model was + * trained on. Several of them are genuinely surprising, so they are constants here with citations + * rather than assumptions in someone's head: + * + * - CODON ORDER IS "TCAG", NOT ALPHABETICAL. The vocabulary is built as + * `[a+b+c for a in "TCAG" for b in "TCAG" for c in "TCAG"]`, the standard genetic-code table + * order. An ACGT-ordered vocabulary is a perfectly valid-looking permutation of the same 64 + * tokens and would be wrong at every site. THE HANDOFF'S OWN INFERENCE DRIVER GETS THIS WRONG — + * see the next paragraph before "correcting" anything here to match it. + * + * THE DRIVER'S TOKENIZER DISAGREES WITH THE TRAINING TOKENIZER. This is measured, not suspected. + * + * train_transformer_selection.py:74-85 (TRAINING) 64 codons, TCAG order, '-'->64, '?'->65 + * predict_regression_nexus.py:49-80 (INFERENCE) 60 codons, ALPHABETICAL, everything else ->65 + * + * The driver defines its own `CODON_LIST` / `CODON_TO_IDX` at lines 49-55, then redefines + * `get_codon_token` at line 74 — but never redefines `CODON_TO_IDX`, so the winning function looks up + * the 60-codon alphabetical map. It imports only `PhyloAxialTransformer`, + * `compute_mds_coordinates` and `decode_soft_ordinal_lrt` from the training module (line 28), so the + * training tokenizer is never in scope. Result, checked over all 64 codons: + * + * 63 of 64 codons receive a DIFFERENT token at inference than at training + * ATG: train 35 -> inference 14 TTT: train 0 -> inference 59 + * AAA: train 42 -> inference 0 GGG: train 63 -> inference 42 + * TTA (Leucine) and all three stop codons are ABSENT from the driver's list entirely, so they + * collapse to 65 "unknown" — real leucine data is discarded rather than mistokenised. + * The amino-acid stream is nearly intact: 3 of 64 wrong, the stops, because the driver's + * GENETIC_CODE writes them as '_' while AA_LIST contains '*', so they land on 22 instead of 20. + * + * DM3 USES THE TRAINING VOCABULARY, which is what this file pins. A model trained on TCAG-64 has to + * be served TCAG-64; there is no reading under which the driver's map is the right one for these + * weights. Reproducing the driver here would reproduce a bug. + * + * Note what this does NOT invalidate: the ML team's ONNX-vs-PyTorch parity result. That test feeds + * the same tensors to both sides, so it proves the graph was exported faithfully and is unaffected + * by which tokenizer produced those tensors. The bug is invisible to it by construction — which is + * exactly why it survived to here. + * - GAP AND UNKNOWN ARE DIFFERENT TOKENS, and they differ between the two streams: codons use + * 64/65, amino acids use 21/22. They are not interchangeable — `forward()` gates on + * `(c_cent < 64) & (a_cent < 21)`, so a gap counts as invalid at the central site. + * - MDS IS COMPUTED ON THE PADDED MATRIX. `compute_mds_coordinates` is handed the full + * [max_species, max_species] tensor, zeros included, so the padded region participates in the + * double-centring and the coordinates depend on max_species — not just on the real taxa. A port + * that runs MDS on the N real species and pads afterwards gets different numbers. + * - THE THREE PHYLO TENSORS ARE SITE-INVARIANT. dist_matrix, mds_coords and padding_mask are + * computed once and `expand`ed across sites. They are per-alignment, not per-site, which is what + * makes batching every site into one graph run cheap. + * + * Sources, all in the ML team's handoff (axomeme-2.0-viral-handoff/scripts/): + * predict_regression_nexus.py:161-187 compute_mds_coordinates + * predict_regression_nexus.py:1213-1229 tensor allocation, pad tokens, MDS on the padded matrix + * predict_regression_nexus.py:1271-1275 the site-invariant expands + * predict_regression_nexus.py:981-982 window_size / max_species defaults + * train_transformer_selection.py:73-97 codon + amino acid vocabularies + * train_transformer_selection.py:1591 forward(msa_codons, msa_aas, dist_matrix, mds_coords, padding_mask) + * + * This module is deliberately a LEAF: it imports nothing, computes nothing, and holds no model. It + * describes and it validates. Nothing here should ever pull in onnxruntime-web — the runtime is + * loaded on the AxoMEME path only, and the reachability guard in + * src/test/meme-hit-likelihood.test.js exists to keep that true. + */ + +/** Codon vocabulary order. NOT alphabetical — see the header. */ +export const CODON_ORDER = 'TCAG'; + +/** 64 sense+stop codons occupy 0..63; these two are the sentinels. */ +export const CODON_GAP = 64; +export const CODON_UNKNOWN = 65; + +/** `num_tokens=66` is passed to the model constructor: 64 codons + gap + unknown. */ +export const NUM_CODON_TOKENS = 66; + +/** Amino acid vocabulary, index = position in this string. */ +export const AA_LIST = 'ACDEFGHIKLMNPQRSTVWY*-?'; +export const AA_GAP = 21; // AA_LIST.indexOf('-') +export const AA_UNKNOWN = 22; // AA_LIST.indexOf('?') + +/** + * `forward()` treats a species as present at the central site only if + * `(codon < 64) & (aa < 21) & !padded`. So a gap is NOT a valid observation, and these are the + * thresholds — not `<= 64` / `<= 21`. + */ +export const CODON_VALID_BELOW = 64; +export const AA_VALID_BELOW = 21; + +/** `--max_species` default. Tensors are padded to exactly this many rows. */ +export const MAX_SPECIES_DEFAULT = 512; + +/** + * `--window_size` default, and the value this checkpoint was fine-tuned at. The model reads the + * CENTRAL column (`window_size // 2`), so an even window would shift which codon is scored. + */ +export const WINDOW_SIZE_DEFAULT = 1; + +/** `compute_mds_coordinates(..., n_components=4)`. */ +export const MDS_COMPONENTS = 4; + +/** + * `padding_mask` is TRUE for padded rows — the opposite of the "1 = keep" convention used in most + * attention APIs, and a sign flip here silently masks out every real taxon instead of none. + */ +export const PADDING_MASK_TRUE_MEANS_PADDED = true; + +/** + * The five graph inputs, in `forward()` order. + * + * `dims` uses the symbolic names the export declares dynamic (batch, num_species); everything else + * is fixed by the checkpoint. int64 matters: onnxruntime-web wants a BigInt64Array for these, and + * passing Float32Array of the same values is a type error at session.run, not a silent coercion. + */ +export const INPUT_SPEC = Object.freeze([ + Object.freeze({ + name: 'msa_codons', + dtype: 'int64', + dims: ['batch', 'num_species', 'window_size'], + valueRange: [0, CODON_UNKNOWN], + padValue: CODON_UNKNOWN, + siteInvariant: false + }), + Object.freeze({ + name: 'msa_aas', + dtype: 'int64', + dims: ['batch', 'num_species', 'window_size'], + valueRange: [0, AA_UNKNOWN], + padValue: AA_UNKNOWN, + siteInvariant: false + }), + Object.freeze({ + name: 'dist_matrix', + dtype: 'float32', + dims: ['batch', 'num_species', 'num_species'], + siteInvariant: true + }), + Object.freeze({ + name: 'mds_coords', + dtype: 'float32', + dims: ['batch', 'num_species', 'mds_components'], + siteInvariant: true + }), + Object.freeze({ + name: 'padding_mask', + dtype: 'bool', + dims: ['batch', 'num_species'], + siteInvariant: true + }) +]); + +/** + * The five graph outputs, in graph order. RESOLVED against the artifact itself — these names were + * read out of `InferenceSession.outputNames` on axomeme_2.0_viral_finetuned.onnx + * (sha256 3e06b591…6faec6), not from documentation. + * + * THE EXPORT IS EVAL MODE. This was an open question worth recording, because the shipped driver + * switches the module to `model.train()` before every forward pass purely to reach the branch that + * returns RAW ORDINAL LOGITS, so it can apply `--prior_shift` and call `decode_soft_ordinal_lrt` + * itself (predict_regression_nexus.py:1354-1366). The answer is that the graph took the EVAL branch: + * five outputs, `lrt` already decoded. Two consequences for the JS side: + * + * - JS does NOT implement `decode_soft_ordinal_lrt`. The graph did it. + * - `prior_shift` is baked to 0 and is not reachable. If a calibration shift is ever wanted, it has + * to be re-exported, not applied here. + * + * STILL OPEN, and it belongs to POSTPROCESSING rather than to this contract: the driver applies + * `np.expm1` to alpha / beta_neg / beta_pos before writing its CSV. Those heads are `F.softplus(...)`, + * and expm1(softplus(x)) == exp(x), so the model is predicting log1p(rate) and expm1 recovers the + * rate — meaning the raw graph outputs are NOT rates and must not be shown as such. `p_neg` is + * sigmoid in-graph and is already a probability. Whether `lrt` needs a further exp is unresolved: the + * driver emits both `predicted_log_lrt` and `predicted_lrt` columns, and which one this output + * corresponds to has to be checked against sample_predictions.csv before anything is rendered. + */ +export const OUTPUT_SPEC = Object.freeze([ + Object.freeze({ + name: 'lrt', + note: 'MEME LRT surrogate, ordinal decode already applied in-graph (eval-mode export)' + }), + Object.freeze({ name: 'alpha', note: 'dS as log1p(rate); caller applies expm1' }), + Object.freeze({ name: 'beta_neg', note: 'as log1p(rate); caller applies expm1' }), + Object.freeze({ name: 'beta_pos', note: 'dN+ as log1p(rate); caller applies expm1' }), + Object.freeze({ name: 'p_neg', note: 'already sigmoid in-graph; a probability as-is' }) +]); + +/** + * The artifact this contract was verified against. A different export is not necessarily wrong, but + * it is not this one, and the eval-mode conclusion above was read off THIS graph. + */ +export const VERIFIED_MODEL_SHA256 = + '3e06b591a060fca996a41c040c2c29f319aa47ca3d3401f4757571b57e6faec6'; + +/** Every input name, in graph order. */ +export const INPUT_NAMES = Object.freeze(INPUT_SPEC.map((s) => s.name)); + +/** + * Check a bundle of prepared tensors against the contract, without running anything. + * + * This is the cheap half of parity: it cannot tell you the MDS coordinates are RIGHT — only real + * fixtures from the Python pipeline can do that — but it catches the whole class of errors that + * produce a well-formed tensor with the wrong meaning, which is the class a JS port actually + * generates. Returns every problem it finds rather than throwing on the first, because a port under + * development usually has several and stopping at one wastes a round trip. + * + * @param {Record, dims: number[]}>} bundle + * @param {{batch: number, numSpecies: number, windowSize?: number, mdsComponents?: number}} shape + * @returns {{ok: boolean, errors: string[]}} + */ +export function validateInputBundle(bundle, shape) { + const errors = []; + const { batch, numSpecies } = shape; + const windowSize = shape.windowSize ?? WINDOW_SIZE_DEFAULT; + const mdsComponents = shape.mdsComponents ?? MDS_COMPONENTS; + + const expected = { + msa_codons: [batch, numSpecies, windowSize], + msa_aas: [batch, numSpecies, windowSize], + dist_matrix: [batch, numSpecies, numSpecies], + mds_coords: [batch, numSpecies, mdsComponents], + padding_mask: [batch, numSpecies] + }; + + for (const spec of INPUT_SPEC) { + const t = bundle[spec.name]; + if (!t) { + errors.push(`${spec.name}: missing`); + continue; + } + const want = expected[spec.name]; + const got = Array.from(t.dims ?? []); + if (got.length !== want.length || got.some((d, i) => d !== want[i])) { + errors.push(`${spec.name}: dims [${got}], expected [${want}]`); + continue; + } + const want_n = want.reduce((a, b) => a * b, 1); + if (t.data.length !== want_n) { + errors.push( + `${spec.name}: ${t.data.length} elements for dims [${want}] (expected ${want_n})` + ); + continue; + } + if (spec.valueRange) { + const [lo, hi] = spec.valueRange; + for (let i = 0; i < t.data.length; i++) { + const v = Number(t.data[i]); + if (!Number.isInteger(v) || v < lo || v > hi) { + errors.push(`${spec.name}[${i}] = ${v}, outside ${lo}..${hi}`); + break; + } + } + } + } + + // Cross-tensor invariants — the ones that are individually well-formed but jointly wrong. + const pad = bundle.padding_mask; + const dist = bundle.dist_matrix; + if (pad && dist && pad.data.length === batch * numSpecies) { + // THE MASK-POLARITY CHECK. `padding_mask` is TRUE for padded rows, the opposite of the + // "1 = keep" convention most attention APIs use, and an inverted mask is well-formed in every + // other respect — right dtype, right dims, right element count. What gives it away is the + // DISTANCE MATRIX: padded rows are never written, so they stay all-zero, while real rows + // carry real distances. Invert the mask and rows holding genuine distances get marked padded. + // + // An earlier version of this checked only self-distances and only the all-padded case. It did + // not fire on a partial flip, which is the flip that actually happens — some rows stay + // unpadded and every self-distance is legitimately zero either way. + let flagged = false; + for (let b = 0; b < batch && !flagged; b++) { + for (let i = 0; i < numSpecies && !flagged; i++) { + const padded = Boolean(pad.data[b * numSpecies + i]); + if (!padded) continue; + const row = b * numSpecies * numSpecies + i * numSpecies; + for (let j = 0; j < numSpecies; j++) { + const v = Number(dist.data[row + j]); + if (v !== 0) { + errors.push( + `padding_mask: species ${i} is marked padded but has distance ${v} to species ` + + `${j} — padded rows are never written, so this is a mask polarity error ` + + '(the convention is TRUE = PADDED)' + ); + flagged = true; + break; + } + } + } + } + const real = Array.from(pad.data.slice(0, numSpecies)).filter((v) => !v).length; + if (real === 0) { + errors.push('padding_mask: every species is marked padded — no taxa would reach the model'); + } + } + + if (dist) { + // d(i,i) = 0 by definition. A nonzero diagonal means whatever was built is not a distance + // matrix — most often an adjacency or a similarity matrix that took the same code path. + for (let b = 0; b < batch; b++) { + for (let i = 0; i < numSpecies; i++) { + const self = Number(dist.data[b * numSpecies * numSpecies + i * numSpecies + i]); + if (self !== 0) { + errors.push(`dist_matrix: self-distance d(${i},${i}) is ${self}, expected 0`); + b = batch; + break; + } + } + } + } + + if (dist) { + for (let i = 0; i < dist.data.length; i++) { + const v = Number(dist.data[i]); + if (!Number.isFinite(v)) { + errors.push(`dist_matrix[${i}] is ${v}`); + break; + } + if (v < 0) { + // Not a shape error — a real one. DM3's own NJ emits negative branch lengths and the + // Python inference path throws on them rather than degrading; see + // src/lib/utils/treeSanitation.js for the measurement and the crash site. + errors.push(`dist_matrix[${i}] = ${v} — negative patristic distance`); + break; + } + } + } + + return { ok: errors.length === 0, errors }; +} diff --git a/src/lib/services/axomeme/newick.js b/src/lib/services/axomeme/newick.js new file mode 100644 index 0000000..1167a6e --- /dev/null +++ b/src/lib/services/axomeme/newick.js @@ -0,0 +1,222 @@ +/** + * newick.js — a minimal newick parser for the AxoMEME preprocessing path. + * + * WHY NOT phylotree, WHICH THIS REPO ALREADY DEPENDS ON. phylotree is a rendering library: it pulls + * d3, it is built for interactive branch selection, and it is used in exactly one place + * (BranchSelector.svelte). What is needed here is the opposite of that — a leaf module with no + * dependencies that produces a tree whose branch lengths and names match, exactly, what Biopython's + * `Phylo.read` hands to the Python reference implementation. Parity is the whole job, and it is much + * easier to defend against a 100-line parser whose semantics are written down than against a + * general-purpose library that also happens to draw SVG. + * + * WHAT IT DELIBERATELY MATCHES (predict_regression_nexus.py:895-905): + * - A MISSING BRANCH LENGTH IS 0.0, not null and not NaN. The reference walks with + * `current_dist + (child.branch_length or 0.0)`, so an absent length contributes nothing rather + * than poisoning every descendant. Note this also maps a length of exactly 0.0 to 0.0, so the + * `or` is harmless there — but a NEGATIVE length is preserved, which matters: DM3's own NJ emits + * them (see src/lib/utils/treeSanitation.js) and silently clamping here would hide that. + * - INTERNAL NODES KEEP THEIR LABELS but are never leaves. A bootstrap value like `)95:0.05` is a + * label on an internal node, and reading it as a taxon name would invent a species. + * + * WHAT IT DELIBERATELY DOES NOT DO: it does not clamp, normalise, or repair anything. A tree that is + * wrong arrives wrong, so the caller can refuse it with the user told rather than scoring a quietly + * altered tree. + */ + +/** + * Strip a quoted or bare newick label down to the name the reference pipeline compares against. + * + * The reference normalises with `s.replace("'", "").replace('"', '').strip()` — Python's str.replace + * removes EVERY occurrence, not just the first — and does it at comparison time + * (predict_regression_nexus.py:1221-1222). Quotes are removed wherever they appear, not just at the + * ends, so a label like `Homo_'sapiens'` and `Homo_sapiens` normalise to the same taxon. + * + * @param {string} name + * @returns {string} + */ +export function normalizeTaxonName(name) { + return String(name).replaceAll("'", '').replaceAll('"', '').trim(); +} + +/** + * Parse a newick string. + * + * Returns a flat node array plus the root index, rather than a linked object graph. That shape is + * chosen on purpose: every consumer here walks the tree by index (root distances, LCA, Max-PD), flat + * typed arrays keep those walks allocation-free, and there are no parent cycles for a structured + * clone or a test's `toEqual` to trip over. + * + * @param {string} text + * @returns {{ + * name: string[], parent: Int32Array, branchLength: Float64Array, + * children: number[][], depth: Int32Array, root: number, leaves: number[] + * }} + */ +export function parseNewick(text) { + if (typeof text !== 'string' || !text.trim()) { + throw new Error('parseNewick: empty input'); + } + + // Strip newick comments before tokenising. `[&&NHX:...]` and `[100]` both appear in the wild and + // neither is part of the topology; leaving them in makes a comment body look like a label. + const src = text.replace(/\[[^\]]*\]/g, ''); + + const name = []; + const parent = []; + const branchLength = []; + const children = []; + + const newNode = (parentIdx) => { + const idx = name.length; + name.push(''); + parent.push(parentIdx); + branchLength.push(0); + children.push([]); + if (parentIdx >= 0) children[parentIdx].push(idx); + return idx; + }; + + let i = 0; + const n = src.length; + const skipSpace = () => { + while (i < n && /\s/.test(src[i])) i++; + }; + + /** Read a label: quoted (quotes preserved, caller normalises) or bare up to a delimiter. */ + const readLabel = () => { + skipSpace(); + if (i >= n) return ''; + const q = src[i]; + if (q === "'" || q === '"') { + // A quoted label may legally contain ':' ',' '(' ')' — which is the entire reason quoting + // exists, and the reason a naive split on ':' corrupts these names. + let out = q; + i++; + while (i < n && src[i] !== q) out += src[i++]; + if (i < n) out += src[i++]; // closing quote + return out; + } + let out = ''; + while (i < n && !'(),:;'.includes(src[i])) out += src[i++]; + return out.trim(); + }; + + /** Read `:length` if present. Absent means 0, matching `branch_length or 0.0`. */ + const readLength = (node) => { + skipSpace(); + if (src[i] !== ':') return; + i++; + skipSpace(); + let out = ''; + while (i < n && !'(),;'.includes(src[i])) out += src[i++]; + const v = parseFloat(out); + // A malformed length is 0, not NaN. NaN would propagate silently through every descendant's + // root distance and surface much later as an all-NaN model input. + branchLength[node] = Number.isFinite(v) ? v : 0; + }; + + // Iterative descent. A recursive parser is simpler but a 5,000-taxon ladder tree is 5,000 frames + // deep, and the reference implementation has exactly that latent bug (Python's traverse() is + // recursive and caps at ~1,000 frames). + const root = newNode(-1); + let current = root; + skipSpace(); + + if (src[i] === '(') { + i++; + current = newNode(root); + while (i < n) { + skipSpace(); + const c = src[i]; + if (c === '(') { + i++; + current = newNode(current); + } else if (c === ',') { + i++; + current = newNode(parent[current]); + } else if (c === ')') { + i++; + current = parent[current]; + name[current] = readLabel(); // internal label / bootstrap + readLength(current); + if (current === root) break; + } else if (c === ';' || i >= n) { + break; + } else { + name[current] = readLabel(); + readLength(current); + } + } + } else { + // A bare single-taxon "tree". + name[root] = readLabel(); + readLength(root); + } + + const count = name.length; + const depth = new Int32Array(count); + const leaves = []; + // PREORDER, depth-first, left to right — matching Biopython's `tree.get_terminals()`, which is + // `find_clades(terminal=True, order='preorder')`. The order is not cosmetic: the reference builds + // its name lookup as `{leaf.name: leaf for leaf in leaves}`, so on a tree with duplicate tip names + // the LAST leaf in THIS order wins, and 1.1% of real DM3 trees have duplicate tips. A + // breadth-first traversal produces the same set and a different winner. + const stack = [root]; + depth[root] = 0; + while (stack.length) { + const node = stack.pop(); + if (children[node].length === 0) { + leaves.push(node); + continue; + } + // Reversed, so the leftmost child is popped first and the walk is genuinely left-to-right. + for (let k = children[node].length - 1; k >= 0; k--) { + const c = children[node][k]; + depth[c] = depth[node] + 1; + stack.push(c); + } + } + + return { + name, + parent: Int32Array.from(parent), + branchLength: Float64Array.from(branchLength), + children, + depth, + root, + leaves + }; +} + +/** + * Map normalised taxon name -> node index, for the named leaves only. + * + * A DUPLICATE NAME RESOLVES TO THE LAST LEAF IN PREORDER, matching the reference's + * `{leaf.name: leaf for leaf in leaves if leaf.name}` — a dict comprehension, so later entries + * overwrite earlier ones. + * + * This was originally written the other way round, keeping the first leaf, on the reasoning that a + * tree with two identically-named tips is malformed and reproducing a dict-overwrite artifact was + * not worth doing. Measurement settled it: 3 of 270 real DM3 trees (1.1%) have duplicate tips — one + * of them has only 36 distinct names across 100 tips — and on every one of those the two policies + * pick different tips and therefore compute different distances. The parity run failed on exactly + * those three and nothing else. + * + * So this layer matches the reference, because its whole job is fidelity to what the model was + * trained and served on, and `duplicates` is returned so the CALLER can refuse the tree. Refusing is + * a product decision; quietly disagreeing with the reference is not. + * + * @param {ReturnType} tree + * @returns {{index: Map, duplicates: string[]}} + */ +export function leafIndex(tree) { + const index = new Map(); + const duplicates = []; + for (const leaf of tree.leaves) { + const nm = normalizeTaxonName(tree.name[leaf]); + if (!nm) continue; + if (index.has(nm)) duplicates.push(nm); + index.set(nm, leaf); // last wins + } + return { index, duplicates }; +} diff --git a/src/lib/services/axomeme/patristic.js b/src/lib/services/axomeme/patristic.js new file mode 100644 index 0000000..e639da2 --- /dev/null +++ b/src/lib/services/axomeme/patristic.js @@ -0,0 +1,185 @@ +/** + * patristic.js — pairwise tree distances and Max-PD taxon selection. + * + * This is the first half of the AxoMEME preprocessing port, and the half that CAN be made + * bit-identical to the Python reference. It is plain float arithmetic over a tree walk: no + * eigendecomposition, no library-version-dependent LAPACK, nothing whose result is only unique up to + * a sign. That is why it comes first — it establishes the parity harness on the part of the pipeline + * where a mismatch is unambiguously a bug rather than a convention difference. + * + * THE REFERENCE (predict_regression_nexus.py:895-965, 1164-1188): + * d(a, b) = rootDist(a) + rootDist(b) - 2 * rootDist(lca(a, b)) + * computed over root distances accumulated as `current_dist + (child.branch_length or 0.0)`. + * + * ONE DELIBERATE DIVERGENCE, AND IT CHANGES NO RESULT. The reference materialises the full N x N + * matrix before Max-PD (`D_full = np.zeros((N_full, N_full))`). Max-PD is a farthest-point traversal + * that only ever reads ONE ROW per iteration, so this port computes rows on demand: `max_species + 1` + * rows instead of N^2 cells. For a 5,000-taxon submission that is ~10 MB instead of ~100 MB, which is + * the difference between running in a browser tab and not. The selected taxa are identical because + * the values are identical — only the cells that are never read go uncomputed. + * + * NEGATIVE DISTANCES ARE PRESERVED, NOT CLAMPED. DM3's own NJ emits negative branch lengths and 5% of + * real DM3 trees carry one at or past -0.1; the Python inference path throws on those rather than + * degrading. Clamping here would convert a loud failure into a quiet wrong answer. See + * src/lib/utils/treeSanitation.js for the measurement and the crash site. + */ + +/** + * Distance from the root to every node, accumulated down the tree. + * + * @param {import('./newick.js').default | any} tree from parseNewick + * @returns {Float64Array} + */ +export function rootDistances(tree) { + const dist = new Float64Array(tree.name.length); + // Explicit stack rather than recursion: a ladder-shaped tree is as deep as it is wide, and the + // Python reference recurses (so it caps out around 1,000 taxa on such a tree). + const stack = [tree.root]; + dist[tree.root] = 0; + while (stack.length) { + const node = stack.pop(); + for (const c of tree.children[node]) { + dist[c] = dist[node] + tree.branchLength[c]; + stack.push(c); + } + } + return dist; +} + +/** + * A reusable scratch buffer for ancestor marking. + * + * LCA is found by marking every ancestor of `a` and walking up from `b` to the first mark. The naive + * version clears the mark array per pair, which is O(nodes) per pair and dominates the whole + * computation for a wide tree. Stamping with a monotonically increasing generation makes the clear + * free. + */ +function ancestorMarker(nodeCount) { + const stamp = new Int32Array(nodeCount); + let generation = 0; + return { + /** Mark the root path of `node`, then return an `isAncestor` predicate for it. */ + markPath(tree, node) { + generation++; + let cur = node; + while (cur !== -1) { + stamp[cur] = generation; + cur = tree.parent[cur]; + } + }, + isMarked(node) { + return stamp[node] === generation; + } + }; +} + +/** + * One row of the patristic distance matrix: distances from `fromNode` to each of `toNodes`. + * + * @param {any} tree + * @param {Float64Array} rootDist from rootDistances() + * @param {number} fromNode + * @param {number[]} toNodes node indices, in the order the row should come out + * @param {ReturnType} [marker] reused across rows by patristicMatrix/maxPd + * @returns {Float64Array} + */ +export function patristicRow(tree, rootDist, fromNode, toNodes, marker) { + const m = marker ?? ancestorMarker(tree.name.length); + m.markPath(tree, fromNode); + const out = new Float64Array(toNodes.length); + for (let k = 0; k < toNodes.length; k++) { + const b = toNodes[k]; + if (b === fromNode) { + out[k] = 0; + continue; + } + let cur = b; + while (cur !== -1 && !m.isMarked(cur)) cur = tree.parent[cur]; + // cur === -1 cannot happen for two nodes of the same tree (the root is always marked), but a + // caller can pass a node from a different tree, and 0 is a less destructive answer than NaN. + out[k] = cur === -1 ? 0 : rootDist[fromNode] + rootDist[b] - 2 * rootDist[cur]; + } + return out; +} + +/** + * The full N x N patristic matrix, row-major. + * + * Only for N small enough to want it whole — the model input is [max_species, max_species], so this + * is the right call there. Max-PD deliberately does NOT use it; see maxPdSelect. + * + * @param {any} tree + * @param {number[]} nodes node indices, defining row and column order + * @returns {Float64Array} length nodes.length ** 2 + */ +export function patristicMatrix(tree, nodes) { + const n = nodes.length; + const rootDist = rootDistances(tree); + const marker = ancestorMarker(tree.name.length); + const out = new Float64Array(n * n); + for (let i = 0; i < n; i++) { + const row = patristicRow(tree, rootDist, nodes[i], nodes, marker); + out.set(row, i * n); + } + return out; +} + +/** + * Max-PD (Faith's PD) farthest-point traversal, matching predict_regression_nexus.py:1174-1183. + * + * selected = [0] + * minDist = D[0] + * repeat: next = argmax(minDist); selected.push(next); minDist = min(minDist, D[next]) + * + * TWO BEHAVIOURS REPRODUCED ON PURPOSE, BOTH OF WHICH LOOK LIKE BUGS: + * + * 1. THE SEED IS ALWAYS INDEX 0 — the first taxon in alignment order, not the most divergent one + * and not a deterministic function of the tree. Reordering the sequences in an upload changes + * which 512 taxa the model sees. That is what the model was fine-tuned against, so this port + * matches it; it is not a defect this layer gets to fix. + * + * 2. A SELECTED TAXON CAN REPEAT. Every selected index has minDist 0 (its own distance to itself + * enters the running minimum), so once every remaining candidate is also at 0 — an all-zero + * distance matrix, i.e. a tree whose branch lengths are all zero — argmax returns index 0 over + * and over and the same taxon fills every slot. `duplicates` reports it rather than silently + * deduplicating, because deduplicating would change which taxa reach a model that was trained + * with the duplicates present. + * + * @param {any} tree + * @param {number[]} nodes candidate node indices, IN ALIGNMENT ORDER (index 0 is the seed) + * @param {number} maxSpecies + * @returns {{selected: number[], duplicates: number}} selected are indices INTO `nodes` + */ +export function maxPdSelect(tree, nodes, maxSpecies) { + const n = nodes.length; + if (n <= maxSpecies) { + return { selected: Array.from({ length: n }, (_, i) => i), duplicates: 0 }; + } + const rootDist = rootDistances(tree); + const marker = ancestorMarker(tree.name.length); + + const selected = [0]; + const seen = new Set([0]); + let duplicates = 0; + // One row, reused. This is the whole memory argument in the header: never N x N, only N. + const minDist = patristicRow(tree, rootDist, nodes[0], nodes, marker); + + for (let step = 1; step < maxSpecies; step++) { + // argmax with FIRST-max-wins on ties, matching np.argmax. + let best = 0; + let bestVal = minDist[0]; + for (let k = 1; k < n; k++) { + if (minDist[k] > bestVal) { + bestVal = minDist[k]; + best = k; + } + } + selected.push(best); + if (seen.has(best)) duplicates++; + else seen.add(best); + + const row = patristicRow(tree, rootDist, nodes[best], nodes, marker); + for (let k = 0; k < n; k++) if (row[k] < minDist[k]) minDist[k] = row[k]; + } + return { selected, duplicates }; +} diff --git a/src/lib/services/axomeme/postprocess.js b/src/lib/services/axomeme/postprocess.js new file mode 100644 index 0000000..cbf50a5 --- /dev/null +++ b/src/lib/services/axomeme/postprocess.js @@ -0,0 +1,249 @@ +/** + * postprocess.js — turning the graph's five raw output tensors into per-site results. + * + * THE RAW OUTPUTS ARE NOT THE NUMBERS A USER SHOULD SEE, and the transformations are not cosmetic: + * + * - `alpha` / `beta_pos` are in LOG1P SPACE. The heads are `F.softplus(...)`, and + * expm1(softplus(x)) == exp(x), so the model predicts log1p(rate) and expm1 recovers the rate. + * Rendering them raw understates every rate — a dN of 1.7 shows as 1.0. + * - `p_neg` is a probability of NEGATIVE selection; the reported quantity is p_pos = 1 - p_neg. + * - `lrt` IS the LRT, already ordinal-decoded in-graph (the export is eval mode). It is NOT a log, + * which was worth checking rather than assuming: the reference derives + * `predicted_log_lrt = log1p(predicted_lrt)`, so the log column is the DERIVED one. + * - `beta_neg` is computed by the model and then never used by the reference's output. Carried here + * so it is available, but it appears in no column and no call. + * + * INVARIANT SITES ARE ZEROED, NOT SCORED. `if not is_var:` sets every prediction to 0 before the + * model's output is consulted (predict_regression_nexus.py:1394-1399). A site where every sequence + * codes the same amino acid reports 0, whatever the network said. This is the reference's behaviour + * and it matters for the UI: those zeros are "not applicable", not "no selection", and they are + * excluded from the z-score and percentile statistics for exactly that reason. + * + * Source: predict_regression_nexus.py:1386-1466. + */ + +import { GENETIC_CODE, aaToken } from './tokenizer.js'; + +/** + * Default tier gates. + * + * THE DEFAULT MODE IS `percentile`, WHICH IS NOT THE REFERENCE'S DEFAULT, and the reason is measured + * rather than preferential. The reference defaults to `pvalue`, which compares the predicted LRT + * against 4.45 and 3.12 — chi-square thresholds for a GENUINE likelihood ratio. The model's output + * does not live on that scale. Across 12 real DataMonkey submissions and 662 variable sites, the + * highest predicted LRT anywhere was 3.902; exactly one site cleared 3.12 and none cleared 4.45. + * On an alignment where MEME itself reports 17 sites at p <= 0.05, the model's maximum was 2.484. + * + * So under `pvalue` this feature ships reporting nothing on real data. That is not a threshold worth + * tuning — it reflects what the model is: a RANKER. The metric its authors report is Spearman rank + * correlation, not calibration, and rank correlation can be good while absolute scale is off. + * `percentile` asks the question the model can answer — which sites in THIS alignment look most + * interesting — instead of one it cannot. + * + * The gates themselves are unchanged from the driver's argparse (lines 984-991), so switching modes + * reproduces the reference exactly. + */ +export const CALL_DEFAULTS = Object.freeze({ + mode: 'percentile', + tier1LrtGate: 4.45, // p <= 0.05 + tier2LrtGate: 3.12, // p <= 0.10 + tier1Zscore: 2.5, + tier2Zscore: 2.0, + tier1Percentile: 98.0, + tier2Percentile: 95.0 +}); + +export const NEUTRAL_CALL = 'Neutral'; + +/** + * Is this site variable, in the reference's sense? + * + * Two conditions, and the second is easy to miss: a site is also variable if every sequence codes + * SERINE but reaches it through both codon families (TCN and AGY). Serine is the one residue whose + * codons occupy two disjoint blocks of the genetic code, so a TCN<->AGY switch requires multiple + * substitutions while remaining synonymous — selection-relevant despite the amino acid never + * changing. Dropping this condition silently marks those sites invariant and zeroes them. + * + * @param {string[]} codons observed codons at this site, gaps/ambiguity already excluded + * @returns {boolean} + */ +export function isSiteVariable(codons) { + if (!codons || codons.length === 0) return false; + const aas = new Set(); + for (const c of codons) { + const aa = GENETIC_CODE.get(c.toUpperCase()); + if (aa && aa !== '?') aas.add(aa); + } + if (aas.size === 0) return false; + if (aas.size > 1) return true; + if (aas.size === 1 && aas.has('S')) { + const upper = codons.map((c) => c.toUpperCase()); + const hasTCN = upper.some((c) => c === 'TCA' || c === 'TCC' || c === 'TCG' || c === 'TCT'); + const hasAGY = upper.some((c) => c === 'AGC' || c === 'AGT'); + if (hasTCN && hasAGY) return true; + } + return false; +} + +/** Percentile rank, matching pandas `rank(pct=True) * 100` — average rank for ties. */ +function percentileRanks(values) { + const n = values.length; + const order = Array.from({ length: n }, (_, i) => i).sort((a, b) => values[a] - values[b]); + const ranks = new Float64Array(n); + let i = 0; + while (i < n) { + let j = i; + while (j + 1 < n && values[order[j + 1]] === values[order[i]]) j++; + // pandas' default tie method is 'average': tied entries share the mean of their 1-based ranks. + const avg = (i + 1 + (j + 1)) / 2; + for (let k = i; k <= j; k++) ranks[order[k]] = avg; + i = j + 1; + } + for (let k = 0; k < n; k++) ranks[k] = (ranks[k] / n) * 100; + return ranks; +} + +/** + * Build the per-site result table. + * + * @param {{lrt: ArrayLike, alpha: ArrayLike, beta_neg: ArrayLike, + * beta_pos: ArrayLike, p_neg: ArrayLike}} outputs raw graph outputs, one per site + * @param {{refCodons: string[], variable: boolean[]}} sites reference codon and variability per site + * @param {object} [callOptions] overrides for CALL_DEFAULTS + * @returns {Array} one row per codon site, 1-indexed `site` + */ +export function buildPredictions(outputs, sites, callOptions = {}) { + const cfg = { ...CALL_DEFAULTS, ...callOptions }; + const n = sites.refCodons.length; + const rows = []; + + for (let i = 0; i < n; i++) { + const isVar = Boolean(sites.variable[i]); + const refCodon = (sites.refCodons[i] ?? '').toUpperCase(); + // translate_codon() rejects gaps, N and '?' before the table lookup, so a gapped reference + // codon has no amino acid rather than an accidental one. + const refAa = + refCodon.length === 3 && + !refCodon.includes('-') && + !refCodon.includes('N') && + !refCodon.includes('?') + ? (GENETIC_CODE.get(refCodon) ?? '?') + : '?'; + + if (!isVar) { + // Everything zeroed BEFORE the model is consulted. These zeros mean "not applicable". + rows.push({ + site: i + 1, + refCodon, + refAa, + isVariable: false, + lrt: 0, + logLrt: 0, + alphaDs: 0, + betaPosDn: 0, + pPos: 0, + zScore: 0, + percentile: 0, + call: NEUTRAL_CALL + }); + continue; + } + + const lrt = Math.max(0, Number(outputs.lrt[i])); + rows.push({ + site: i + 1, + refCodon, + refAa, + isVariable: true, + lrt, + logLrt: Math.log1p(lrt), + // expm1 undoes the softplus/log1p parameterisation; the clamp is the reference's. + alphaDs: Math.max(0, Math.expm1(Number(outputs.alpha[i]))), + betaPosDn: Math.max(0, Math.expm1(Number(outputs.beta_pos[i]))), + pPos: 1 - Number(outputs.p_neg[i]), + zScore: 0, + percentile: 0, + call: NEUTRAL_CALL + }); + } + + // Local statistics are computed over VARIABLE SITES ONLY. Including the zeroed invariant sites + // would drag the mean down and inflate every z-score, which is the whole reason they are excluded. + const varIdx = rows.map((r, i) => (r.isVariable ? i : -1)).filter((i) => i >= 0); + if (varIdx.length === 0) return rows; + + const varLrts = varIdx.map((i) => rows[i].lrt); + const mean = varLrts.reduce((a, b) => a + b, 0) / varLrts.length; + // Population standard deviation — np.std, not pandas' sample std. + const variance = varLrts.reduce((a, b) => a + (b - mean) ** 2, 0) / varLrts.length; + const std = Math.sqrt(variance); + + const pct = percentileRanks(varLrts); + varIdx.forEach((rowIdx, k) => { + rows[rowIdx].zScore = std > 0 ? (varLrts[k] - mean) / std : 0; + rows[rowIdx].percentile = pct[k]; + }); + + const tierLabels = + cfg.mode === 'zscore' + ? { tier1: `Z \u2265 ${cfg.tier1Zscore}`, tier2: `Z \u2265 ${cfg.tier2Zscore}` } + : cfg.mode === 'pvalue' + ? { tier1: `LRT \u2265 ${cfg.tier1LrtGate}`, tier2: `LRT \u2265 ${cfg.tier2LrtGate}` } + : { + tier1: `Top ${(100 - cfg.tier1Percentile).toFixed(0)}%`, + tier2: `Top ${(100 - cfg.tier2Percentile).toFixed(0)}%` + }; + + for (const i of varIdx) { + const r = rows[i]; + let t1 = false; + let t2 = false; + if (cfg.mode === 'zscore') { + t1 = r.zScore >= cfg.tier1Zscore; + t2 = !t1 && r.zScore >= cfg.tier2Zscore; + } else if (cfg.mode === 'percentile') { + t1 = r.percentile >= cfg.tier1Percentile; + t2 = !t1 && r.percentile >= cfg.tier2Percentile; + } else { + t1 = r.lrt >= cfg.tier1LrtGate; + t2 = !t1 && r.lrt >= cfg.tier2LrtGate; + } + // Labels say what the tier MEANS rather than how confident it sounds. "High" and "Medium" + // imply a calibrated confidence the model does not have; "Top 2%" is exactly what percentile + // mode computed, and a reader can tell at a glance that it is relative to this alignment. + if (t1) r.call = tierLabels.tier1; + else if (t2) r.call = tierLabels.tier2; + } + + return rows; +} + +/** + * Variability flags for every site of an alignment, from the aligned sequences. + * + * Mirrors the reference's collection loop: only codons that are complete, ungapped and unambiguous + * are considered, and a sequence shorter than the reference simply contributes nothing at the sites + * it does not reach. + * + * @param {string[]} sequences aligned nucleotide sequences, same frame + * @param {number} totalCodons + * @returns {boolean[]} + */ +export function siteVariability(sequences, totalCodons) { + const flags = new Array(totalCodons).fill(false); + for (let s = 0; s < totalCodons; s++) { + const codons = []; + for (const seq of sequences) { + const start = s * 3; + if (start + 3 > seq.length) continue; + const c = seq.slice(start, start + 3).toUpperCase(); + if (c.includes('-') || c.includes('N') || c.includes('?')) continue; + // aaToken rejects anything the genetic code cannot translate, which is the same filter the + // reference applies before collecting a codon. + if (aaToken(c) > 20) continue; + codons.push(c); + } + flags[s] = isSiteVariable(codons); + } + return flags; +} diff --git a/src/lib/services/axomeme/session.js b/src/lib/services/axomeme/session.js new file mode 100644 index 0000000..23f9395 --- /dev/null +++ b/src/lib/services/axomeme/session.js @@ -0,0 +1,178 @@ +/** + * session.js — loading and running the AxoMEME 2.0 ONNX graph in the browser. + * + * THE COST DISCIPLINE THIS FILE EXISTS TO ENFORCE. onnxruntime-web is ~13 MB of WASM and the model is + * another 3.78 MB. DataMonkey has fifteen analysis methods and AxoMEME is one of them; the other + * fourteen must download none of it. That is not a nice-to-have — the first version of the MEME + * hit-likelihood gate shipped 13.5 MB of ONNX Runtime to every method and rendered nothing for + * fourteen of them, and an "is the element absent?" test passed the whole time because the element + * WAS absent. The bytes were the bug. + * + * So: + * - The runtime is loaded by DYNAMIC import inside loadSession(), never at module scope. Importing + * this file costs nothing; calling loadSession() is what costs 17 MB. + * - The model is fetched from static/ rather than bundled, so it is cached by the browser + * independently of the JS and never enters the bundle graph. + * - The session is memoised, because a second 17 MB download to score a second alignment would be + * indefensible. + * + * INTEGRITY. The artifact's sha256 is pinned in modelContract.js and verified after fetch. The + * contract's conclusions — eval-mode export, `lrt` already decoded, rate heads in log1p space — were + * read off THAT graph, and a silently swapped model would make all of them wrong while everything + * continued to run. Verification is skipped only when Web Crypto is unavailable (non-secure context), + * which is reported rather than hidden. + * + * BATCHING. dist_matrix, mds_coords and padding_mask are per-alignment, not per-site — the reference + * computes them once and `expand`s them across sites. So every codon site of an alignment goes + * through the graph in ONE run, with those three tensors repeated along the batch axis. The reference + * driver instead loops one forward pass per site (predict_regression_nexus.py:1345), which for a + * 441-site alignment is 441 sessions of overhead for identical work. + */ + +import { INPUT_NAMES, VERIFIED_MODEL_SHA256 } from './modelContract.js'; + +/** Where the artifact lives. Served from static/, never imported, never bundled. */ +export const MODEL_URL = '/models/axomeme/axomeme_2.0_viral_finetuned.onnx'; + +/** Memoised session promise. Null until the first scoring call. */ +let sessionPromise = null; + +/** + * Hex sha256 of a buffer, or null where Web Crypto is unavailable. + * + * Digests a Uint8Array VIEW rather than the ArrayBuffer itself. `digest` accepts either per spec, but + * an ArrayBuffer constructed in another realm fails the implementation's instanceof check — which is + * exactly what happens under jsdom, and would equally happen for a buffer crossing a worker boundary. + * Constructing the view here puts it in this module's realm. + */ +async function sha256Hex(buffer) { + if (typeof globalThis.crypto?.subtle?.digest !== 'function') return null; + const digest = await globalThis.crypto.subtle.digest('SHA-256', new Uint8Array(buffer)); + return Array.from(new Uint8Array(digest)) + .map((b) => b.toString(16).padStart(2, '0')) + .join(''); +} + +/** + * Load the ONNX session, downloading the runtime and the model on first call. + * + * @param {{url?: string, verifyHash?: boolean, ort?: any, fetchImpl?: typeof fetch}} [options] + * `ort` and `fetchImpl` exist for tests; production passes neither. + * @returns {Promise<{session: any, sha256: string|null, bytes: number}>} + */ +export function loadSession(options = {}) { + // Memoise ONLY the production call — an options object of any kind bypasses the cache. + // + // This used to list url/ort/fetchImpl by name, which left `verifyHash` and `wasmPaths` out: a + // single `loadSession({ verifyHash: false })` would populate sessionPromise with a session that was + // never checked against VERIFIED_MODEL_SHA256, and every later production call would return it + // with `sha256: null` — which the runner then stamps as the verified hash. Inverting the test to + // "no options at all" cannot go stale when a new option is added. + const isProductionCall = Object.keys(options).length === 0; + if (sessionPromise && isProductionCall) return sessionPromise; + + const promise = (async () => { + const url = options.url ?? MODEL_URL; + const doFetch = options.fetchImpl ?? globalThis.fetch; + + // The dynamic import is the whole point — see the header. Do not hoist it. + // + // NOTE THE SUBPATH: 'onnxruntime-web/wasm', not 'onnxruntime-web'. The default entry is the + // full build, which resolves its binary to the JSEP variant — the WebGPU-enabled one, a 26.8 MB + // wasm this feature has no use for. Importing the CPU-only build makes it load the plain + // SIMD+threads binary (12.9 MB) that scripts/copy-ort-wasm.mjs vendors. Found the hard way: + // with the default entry the runtime asks for ort-wasm-simd-threaded.jsep.mjs and aborts with + // "no available backend found", which reads like a broken model rather than a wrong build. + const ort = options.ort ?? (await import('onnxruntime-web/wasm')); + + // SERVE THE RUNTIME OURSELVES. onnxruntime-web does not bundle its WASM binary; it fetches it + // at run time, and with no wasmPaths set it resolves to a jsDelivr CDN. That violates this + // project's core constraint — the site must be servable with no other domains involved — and + // it also simply fails, with "no available backend found", which reads like a broken model + // rather than a missing asset. scripts/copy-ort-wasm.mjs puts the binary in static/ort/ at + // build time, pinned by package.json rather than committed. + // + // numThreads 1 because multi-threading needs SharedArrayBuffer, which needs COOP/COEP headers + // this app does not set. Asking for threads without them makes the runtime fall back anyway, + // and the fallback path is slower than starting single-threaded. + if (ort.env?.wasm) { + ort.env.wasm.wasmPaths = options.wasmPaths ?? '/ort/'; + ort.env.wasm.numThreads = 1; + } + + const response = await doFetch(url); + if (!response.ok) { + throw new Error(`AxoMEME model fetch failed: ${response.status} ${response.statusText}`); + } + const buffer = await response.arrayBuffer(); + + const sha256 = options.verifyHash === false ? null : await sha256Hex(buffer); + if (sha256 && sha256 !== VERIFIED_MODEL_SHA256) { + // Refuse rather than score. Every conclusion in modelContract.js was read off the pinned + // graph; a different one may well be fine, but nothing here knows that, and the failure + // mode of guessing is silently wrong numbers rather than an error. + throw new Error( + `AxoMEME model hash mismatch.\n expected ${VERIFIED_MODEL_SHA256}\n got ${sha256}\n` + + 'The input/output contract was verified against the expected artifact. If the model was ' + + 'deliberately updated, re-verify src/lib/services/axomeme/modelContract.js and update ' + + 'VERIFIED_MODEL_SHA256 in the same commit.' + ); + } + + const session = await ort.InferenceSession.create(new Uint8Array(buffer)); + + // The graph must expose exactly what the contract says. A rename upstream would otherwise + // surface as an opaque runtime error deep inside session.run. + const missing = INPUT_NAMES.filter((n) => !session.inputNames.includes(n)); + if (missing.length) { + throw new Error(`AxoMEME model is missing expected inputs: ${missing.join(', ')}`); + } + + return { session, sha256, bytes: buffer.byteLength }; + })(); + + if (isProductionCall) sessionPromise = promise; + // A failed load must not be memoised as a permanent failure — a transient fetch error would + // otherwise disable the feature for the rest of the page's life. + promise.catch(() => { + if (sessionPromise === promise) sessionPromise = null; + }); + return promise; +} + +/** Drop the memoised session. Tests use this; production has no reason to. */ +export function resetSession() { + sessionPromise = null; +} + +/** True once a session is loaded or loading — lets a caller avoid triggering a 17 MB download. */ +export function isSessionLoaded() { + return sessionPromise !== null; +} + +/** + * Run every site of an alignment through the graph. + * + * @param {any} session an ONNX InferenceSession + * @param {object} bundle tensors as produced by the assembly layer, shaped per modelContract + * @param {any} ort the onnxruntime module (passed in so this stays a pure function) + * @returns {Promise<{lrt: Float32Array, alpha: Float32Array, beta_neg: Float32Array, + * beta_pos: Float32Array, p_neg: Float32Array}>} + */ +export async function runSites(session, bundle, ort) { + const feeds = { + msa_codons: new ort.Tensor('int64', bundle.msa_codons.data, bundle.msa_codons.dims), + msa_aas: new ort.Tensor('int64', bundle.msa_aas.data, bundle.msa_aas.dims), + dist_matrix: new ort.Tensor('float32', bundle.dist_matrix.data, bundle.dist_matrix.dims), + mds_coords: new ort.Tensor('float32', bundle.mds_coords.data, bundle.mds_coords.dims), + padding_mask: new ort.Tensor('bool', bundle.padding_mask.data, bundle.padding_mask.dims) + }; + const out = await session.run(feeds); + return { + lrt: out.lrt.data, + alpha: out.alpha.data, + beta_neg: out.beta_neg.data, + beta_pos: out.beta_pos.data, + p_neg: out.p_neg.data + }; +} diff --git a/src/lib/services/axomeme/symmetricEigen.js b/src/lib/services/axomeme/symmetricEigen.js new file mode 100644 index 0000000..14c7a65 --- /dev/null +++ b/src/lib/services/axomeme/symmetricEigen.js @@ -0,0 +1,227 @@ +/** + * symmetricEigen.js — eigendecomposition of a real symmetric matrix. + * + * WHY THIS IS HERE AT ALL. AxoMEME's `mds_coords` is a model INPUT, not something the ONNX graph + * computes, because `torch.linalg.eigh` has no ONNX lowering — verified directly: the identical + * module exports fine with eigh removed and fails with it present. So the browser has to do its own + * eigendecomposition of the [max_species, max_species] double-centred matrix, which at the default + * max_species is 512 x 512. + * + * THE ALGORITHM is Householder tridiagonalisation followed by implicit-shift QL — `tred2`/`tql2`, + * the EISPACK routines that JAMA and Numerical Recipes both carry, and mathematically what LAPACK's + * symmetric drivers reduce to. It is chosen over Jacobi for cost: Jacobi on 512 x 512 needs roughly + * 10 sweeps of ~131k rotations and lands in the seconds; tred2/tql2 is ~(4/3)n^3 once plus cheap + * iteration, and runs in a few hundred milliseconds. + * + * WHAT PARITY CAN AND CANNOT BE PROMISED HERE, stated plainly because it is the crux of the whole + * browser port: + * + * - EIGENVALUES agree with LAPACK to near machine precision. They are a property of the matrix. + * - EIGENVECTORS are only unique up to sign, and within a DEGENERATE eigenspace not even up to + * sign — any orthonormal basis of that subspace is equally correct. numpy uses divide-and-conquer + * (dsyevd); this is QL. Two correct implementations can therefore disagree, and no amount of care + * in this file changes that. + * + * The sign half is fixed by convention downstream (mds.js applies the reference's own rule). The + * degeneracy half is not fixable in principle, only measurable — which is why the parity harness runs + * over real trees rather than a proof being attempted here. + * + * The routine is transcribed rather than invented, and its tests check the properties that actually + * matter (A·V = V·Λ, orthonormality, reconstruction) rather than comparing against numbers this file + * produced. + */ + +/** Machine epsilon for float64, as the reference routines define it. */ +const EPS = Math.pow(2, -52); + +/** + * Eigendecomposition of a real symmetric matrix. + * + * @param {Float64Array|number[]} matrix row-major, n*n, assumed symmetric (upper triangle unused) + * @param {number} n + * @returns {{values: Float64Array, vectors: Float64Array}} eigenvalues ASCENDING (matching + * numpy.linalg.eigh), eigenvectors row-major with column j being the vector for values[j] + */ +export function symmetricEigen(matrix, n) { + const V = Float64Array.from(matrix); + const d = new Float64Array(n); + const e = new Float64Array(n); + + tred2(V, d, e, n); + tql2(V, d, e, n); + + return { values: d, vectors: V }; +} + +/** + * Householder reduction to tridiagonal form. + * + * On entry V holds the symmetric matrix; on exit V holds the accumulated orthogonal transformation, + * d the diagonal and e the sub-diagonal of the tridiagonal form. + */ +function tred2(V, d, e, n) { + for (let j = 0; j < n; j++) d[j] = V[(n - 1) * n + j]; + + for (let i = n - 1; i > 0; i--) { + let scale = 0; + let h = 0; + for (let k = 0; k < i; k++) scale += Math.abs(d[k]); + + if (scale === 0) { + e[i] = d[i - 1]; + for (let j = 0; j < i; j++) { + d[j] = V[(i - 1) * n + j]; + V[i * n + j] = 0; + V[j * n + i] = 0; + } + } else { + for (let k = 0; k < i; k++) { + d[k] /= scale; + h += d[k] * d[k]; + } + let f = d[i - 1]; + let g = Math.sqrt(h); + if (f > 0) g = -g; + e[i] = scale * g; + h -= f * g; + d[i - 1] = f - g; + for (let j = 0; j < i; j++) e[j] = 0; + + for (let j = 0; j < i; j++) { + f = d[j]; + V[j * n + i] = f; + g = e[j] + V[j * n + j] * f; + for (let k = j + 1; k <= i - 1; k++) { + g += V[k * n + j] * d[k]; + e[k] += V[k * n + j] * f; + } + e[j] = g; + } + + f = 0; + for (let j = 0; j < i; j++) { + e[j] /= h; + f += e[j] * d[j]; + } + const hh = f / (h + h); + for (let j = 0; j < i; j++) e[j] -= hh * d[j]; + + for (let j = 0; j < i; j++) { + f = d[j]; + g = e[j]; + for (let k = j; k <= i - 1; k++) V[k * n + j] -= f * e[k] + g * d[k]; + d[j] = V[(i - 1) * n + j]; + V[i * n + j] = 0; + } + } + d[i] = h; + } + + // Accumulate the transformations. + for (let i = 0; i < n - 1; i++) { + V[(n - 1) * n + i] = V[i * n + i]; + V[i * n + i] = 1; + const h = d[i + 1]; + if (h !== 0) { + for (let k = 0; k <= i; k++) d[k] = V[k * n + (i + 1)] / h; + for (let j = 0; j <= i; j++) { + let g = 0; + for (let k = 0; k <= i; k++) g += V[k * n + (i + 1)] * V[k * n + j]; + for (let k = 0; k <= i; k++) V[k * n + j] -= g * d[k]; + } + } + for (let k = 0; k <= i; k++) V[k * n + (i + 1)] = 0; + } + for (let j = 0; j < n; j++) { + d[j] = V[(n - 1) * n + j]; + V[(n - 1) * n + j] = 0; + } + V[(n - 1) * n + (n - 1)] = 1; + e[0] = 0; +} + +/** Symmetric tridiagonal QL with implicit shifts. Leaves eigenvalues in d, ASCENDING. */ +function tql2(V, d, e, n) { + for (let i = 1; i < n; i++) e[i - 1] = e[i]; + e[n - 1] = 0; + + let f = 0; + let tst1 = 0; + + for (let l = 0; l < n; l++) { + tst1 = Math.max(tst1, Math.abs(d[l]) + Math.abs(e[l])); + let m = l; + while (m < n) { + if (Math.abs(e[m]) <= EPS * tst1) break; + m++; + } + + if (m > l) { + do { + let g = d[l]; + let p = (d[l + 1] - g) / (2 * e[l]); + let r = Math.hypot(p, 1); + if (p < 0) r = -r; + d[l] = e[l] / (p + r); + d[l + 1] = e[l] * (p + r); + const dl1 = d[l + 1]; + let h = g - d[l]; + for (let i = l + 2; i < n; i++) d[i] -= h; + f += h; + + p = d[m]; + let c = 1; + let c2 = c; + let c3 = c; + const el1 = e[l + 1]; + let s = 0; + let s2 = 0; + for (let i = m - 1; i >= l; i--) { + c3 = c2; + c2 = c; + s2 = s; + g = c * e[i]; + h = c * p; + r = Math.hypot(p, e[i]); + e[i + 1] = s * r; + s = e[i] / r; + c = p / r; + p = c * d[i] - s * g; + d[i + 1] = h + s * (c * g + s * d[i]); + for (let k = 0; k < n; k++) { + h = V[k * n + (i + 1)]; + V[k * n + (i + 1)] = s * V[k * n + i] + c * h; + V[k * n + i] = c * V[k * n + i] - s * h; + } + } + p = (-s * s2 * c3 * el1 * e[l]) / dl1; + e[l] = s * p; + d[l] = c * p; + } while (Math.abs(e[l]) > EPS * tst1); + } + d[l] += f; + e[l] = 0; + } + + // Sort ascending, carrying the eigenvectors with their values. Selection sort: n is at most a few + // hundred here and the swap has to move a whole column, so the simple algorithm is the right one. + for (let i = 0; i < n - 1; i++) { + let k = i; + let p = d[i]; + for (let j = i + 1; j < n; j++) { + if (d[j] < p) { + k = j; + p = d[j]; + } + } + if (k !== i) { + d[k] = d[i]; + d[i] = p; + for (let j = 0; j < n; j++) { + const t = V[j * n + i]; + V[j * n + i] = V[j * n + k]; + V[j * n + k] = t; + } + } + } +} diff --git a/src/lib/services/axomeme/tokenizer.js b/src/lib/services/axomeme/tokenizer.js new file mode 100644 index 0000000..749d872 --- /dev/null +++ b/src/lib/services/axomeme/tokenizer.js @@ -0,0 +1,128 @@ +/** + * tokenizer.js — codon and amino-acid tokenisation for AxoMEME 2.0. + * + * THIS IMPLEMENTS THE TRAINING TOKENIZER (train_transformer_selection.py:74-97), NOT THE ONE IN THE + * HANDOFF'S INFERENCE DRIVER. That is a deliberate, measured choice, and it is the single most + * consequential decision in this file — see the long note in modelContract.js. In short: + * predict_regression_nexus.py defines its own 60-codon ALPHABETICAL vocabulary at lines 49-55, then + * redefines `get_codon_token` without redefining `CODON_TO_IDX`, so 63 of 64 codons come out with a + * different token at inference than the model was trained on, and TTA plus the three stop codons + * vanish into "unknown". A model trained on TCAG-64 must be served TCAG-64. + * + * If a future handoff fixes the driver, this file does not change — it already matches training. + * If someone "aligns" this file to the driver, the model silently degrades and nothing fails. + * src/test/axomeme-tokenizer.test.js pins both directions. + */ + +import { + CODON_ORDER, + CODON_GAP, + CODON_UNKNOWN, + AA_LIST, + AA_GAP, + AA_UNKNOWN +} from './modelContract.js'; + +/** + * The 64 codons in TCAG order — the standard genetic-code table order, so TTT is 0 and GGG is 63. + * Built rather than transcribed: a 64-entry literal is a transcription error waiting to happen, and + * the generator IS the reference's definition. + */ +export const CODON_LIST = (() => { + const out = []; + for (const a of CODON_ORDER) + for (const b of CODON_ORDER) for (const c of CODON_ORDER) out.push(a + b + c); + return out; +})(); + +/** codon string -> token, for the 64 sense+stop codons. */ +export const CODON_TO_IDX = new Map(CODON_LIST.map((c, i) => [c, i])); + +/** + * The standard genetic code, with stops as '*' to match AA_LIST. + * + * Derived from the codon table rather than transcribed, for the same reason as CODON_LIST. The four + * blocks below are the standard code's structure; every codon is covered, so there is no default. + * + * ONE THING WORTH FLAGGING UPSTREAM: this is the UNIVERSAL code, hard-coded. DM3 lets a user pick a + * genetic code (mitochondrial, mycoplasma, several others) and HyPhy honours that choice, so an + * AxoMEME prediction on a non-universal alignment translates codons the user did not ask for. That + * is a model-side limitation, not something this file can fix — recorded here so it is not + * rediscovered as a mystery. + */ +export const GENETIC_CODE = (() => { + // TCAG-ordered amino acids, one per codon, in the same order CODON_LIST generates. + const table = + 'FFLLSSSSYY**CC*W' + // TTx TCx TAx TGx + 'LLLLPPPPHHQQRRRR' + // CTx CCx CAx CGx + 'IIIMTTTTNNKKSSRR' + // ATx ACx AAx AGx + 'VVVVAAAADDEEGGGG'; // GTx GCx GAx GGx + const map = new Map(); + CODON_LIST.forEach((c, i) => map.set(c, table[i])); + return map; +})(); + +/** amino acid character -> token. */ +export const AA_TO_IDX = new Map([...AA_LIST].map((a, i) => [a, i])); + +/** + * Token for a codon string. + * + * Order of tests matters and matches the reference exactly: a gap ANYWHERE wins over everything + * else, so `A-T` is a gap rather than an unknown. `'-' in codon` is a substring test in Python, not + * an equality test, which is why a partial gap counts. + * + * @param {string} codon + * @returns {number} 0..63 for a real codon, CODON_GAP for a gap, CODON_UNKNOWN otherwise + */ +export function codonToken(codon) { + const c = String(codon).toUpperCase(); + if (c.includes('-')) return CODON_GAP; + if (c.length !== 3 || c.includes('N')) return CODON_UNKNOWN; + const t = CODON_TO_IDX.get(c); + return t === undefined ? CODON_UNKNOWN : t; +} + +/** + * Token for the amino acid a codon translates to. + * + * Note the asymmetry with codonToken, which is in the reference and is not a typo here: this one + * also rejects '?', codonToken does not. A '?' codon reaches codonToken's length/N test and, being + * length 3 without an N, falls through to the map lookup and returns CODON_UNKNOWN anyway — same + * answer by a different route. + * + * @param {string} codon + * @returns {number} 0..20 for a translated residue (20 = stop), AA_GAP or AA_UNKNOWN + */ +export function aaToken(codon) { + const c = String(codon).toUpperCase(); + if (c.includes('-')) return AA_GAP; + if (c.length !== 3 || c.includes('N') || c.includes('?')) return AA_UNKNOWN; + const aa = GENETIC_CODE.get(c); + if (aa === undefined) return AA_UNKNOWN; + const t = AA_TO_IDX.get(aa); + return t === undefined ? AA_UNKNOWN : t; +} + +/** + * Tokenise one sequence into per-codon (codon, aa) token pairs. + * + * A trailing partial codon is DROPPED, matching `total_codons = len(ref_seq) // 3`. Sites beyond a + * sequence's own length are the caller's problem: the reference leaves those at the pad value it + * pre-filled the tensor with, which is why this returns only what the sequence actually covers. + * + * @param {string} seq nucleotides, gaps allowed + * @returns {{codons: Uint8Array, aas: Uint8Array}} + */ +export function tokenizeSequence(seq) { + const s = String(seq); + const n = Math.floor(s.length / 3); + const codons = new Uint8Array(n); + const aas = new Uint8Array(n); + for (let i = 0; i < n; i++) { + const c = s.slice(i * 3, i * 3 + 3); + codons[i] = codonToken(c); + aas[i] = aaToken(c); + } + return { codons, aas }; +} diff --git a/src/lib/utils/treeSanitation.js b/src/lib/utils/treeSanitation.js new file mode 100644 index 0000000..76369c9 --- /dev/null +++ b/src/lib/utils/treeSanitation.js @@ -0,0 +1,111 @@ +/** + * treeSanitation.js — inspect a newick for branch lengths no downstream consumer should be handed. + * + * WHY THIS EXISTS. DM3's own tree inference emits negative branch lengths, and at least one + * consumer crashes on them rather than degrading. + * + * - src/data/shared/NJ.bf:214-220 computes the three-taxon closed form + * (d01 + d02 - d12) / 2 with no Max(0, ...), so any triplet violating the triangle inequality + * yields a negative. For four or more taxa the NJ core is C++ inside hyphy.wasm and likewise + * does not clamp. + * - src/data/shared/NJ.bf:99 returns a SATURATION SENTINEL of 1000 for a saturated pair, which + * propagates through that same subtraction to roughly -499.9. + * + * Measured consequence, on the AxoMEME 2.0 inference path: + * predict_regression_nexus.py:955 computes log((node_count + 1.0) / (dist + 0.1)). + * Any patristic distance <= -0.1 is log of a negative -> uncaught ValueError. Verified: -0.010 + * survives, -0.11 / -0.5 / -499.9 all crash. The handoff README's "clamps distances >= 0" refers + * to the TRAINING pipeline; the inference path does not clamp. + * + * So this is a DM3-side fact about DM3-side output, and it holds wherever a consumer runs — browser + * or server. It is deliberately a leaf module (imports nothing) so a caller can ask before loading + * anything heavy, matching the pattern in services/prescreen/scope.js. + * + * This module REPORTS. It does not silently rewrite a user's tree: clamping changes branch lengths, + * branch lengths are the input a model's distances are built from, and quietly altering them would + * be a fabrication of the same kind this codebase already refuses elsewhere. The caller decides + * whether to refuse the run or to clamp with the user told. + */ + +/** Matches a newick branch length, including scientific notation and a leading minus. */ +const BRANCH_LENGTH = /:(-?\d+\.?\d*(?:[eE][-+]?\d+)?)/g; + +/** + * The value NJ.bf:99 returns for a saturated pair. It is not a distance; it is a sentinel, and it + * reaches trees as a large positive length or, after the three-taxon subtraction, a large negative. + */ +export const NJ_SATURATION_SENTINEL = 1000; + +/** + * Every branch length in a newick, in document order. Empty array for a topology-only tree. + * @param {string|null} tree + * @returns {number[]} + */ +export function branchLengths(tree) { + if (!tree || typeof tree !== 'string') return []; + const out = []; + let m; + BRANCH_LENGTH.lastIndex = 0; + while ((m = BRANCH_LENGTH.exec(tree)) !== null) { + const v = parseFloat(m[1]); + if (Number.isFinite(v)) out.push(v); + } + return out; +} + +/** + * Describe what is wrong with a tree's branch lengths, without changing it. + * + * @param {string|null} tree + * @returns {{ + * total: number, negative: number, negativeFraction: number, min: number|null, + * saturated: number, hasLengths: boolean, ok: boolean, reasons: string[] + * }} + */ +export function inspectBranchLengths(tree) { + const lengths = branchLengths(tree); + const negatives = lengths.filter((v) => v < 0); + const saturated = lengths.filter((v) => Math.abs(v) >= NJ_SATURATION_SENTINEL); + const reasons = []; + + if (!lengths.length) reasons.push('topology-only: the tree carries no branch lengths'); + if (negatives.length) { + reasons.push( + `${negatives.length} of ${lengths.length} branch lengths are negative ` + + `(smallest ${Math.min(...negatives)})` + ); + } + if (saturated.length) { + reasons.push( + `${saturated.length} branch length(s) at or past the NJ saturation sentinel ` + + `(|value| >= ${NJ_SATURATION_SENTINEL}) — these are not distances` + ); + } + + return { + total: lengths.length, + negative: negatives.length, + negativeFraction: lengths.length ? negatives.length / lengths.length : 0, + min: lengths.length ? Math.min(...lengths) : null, + saturated: saturated.length, + hasLengths: lengths.some((v) => v > 0), + ok: reasons.length === 0, + reasons + }; +} + +/** + * Would a consumer that computes log((n + 1) / (d + 0.1)) over PATRISTIC distances crash? + * + * Deliberately conservative, and the comment matters more than the code: the crash threshold is on + * PATRISTIC distances (root-to-root path sums), not on individual branch lengths, and a path sum can + * be more negative than any single branch on it. A tree can therefore pass this check and still + * crash the consumer. Treat `false` as "no single branch proves a crash", never as "safe". + * + * @param {string|null} tree + * @param {number} [epsilon] the additive constant in the consumer's log argument + * @returns {boolean} + */ +export function hasCrashingBranchLength(tree, epsilon = 0.1) { + return branchLengths(tree).some((v) => v <= -epsilon); +} diff --git a/src/routes/+page.svelte b/src/routes/+page.svelte index 92da3cb..483eca3 100644 --- a/src/routes/+page.svelte +++ b/src/routes/+page.svelte @@ -9,7 +9,12 @@ persistentFileStore, currentFile } from '../stores/fileInfo'; - import { analysisStore, currentAnalysis, activeAnalysisProgress, activeAnalyses } from '../stores/analyses'; + import { + analysisStore, + currentAnalysis, + activeAnalysisProgress, + activeAnalyses + } from '../stores/analyses'; import { backendAnalysisRunner } from '../lib/services/BackendAnalysisRunner.js'; import { treeStore, addTree, updateTaggedTree } from '../stores/tree'; import { trackEvent } from '../lib/utils/analytics.js'; @@ -65,6 +70,23 @@ // Consolidating methods and hyphyCommands with descriptions const methodConfig = { + // AxoMEME is NOT a HyPhy method, and the fields below are shaped for one. It has no hyphy + // command and writes no HyPhy JSON, so `command` and `outputSuffix` are null rather than + // invented — anything that dispatches on them will fail loudly instead of running the wrong + // binary. It is served by AxomemeAnalysisRunner, which evaluates an ONNX graph in the browser. + AxoMEME: { + command: null, + outputSuffix: null, + url: 'axomeme', + args: [], + runner: 'axomeme', + // Drives the UI: no execution-mode toggle (there is no server-side AxoMEME) and no genetic + // code selector (the model's tokenizer bakes in the universal table). A control that cannot + // reach the model is worse than no control, because it implies a capability. + browserOnly: true, + description: + 'Predicts what MEME would report for each site, in seconds rather than hours. A neural surrogate, not a substitute for the full analysis.' + }, aBSREL: { command: 'absrel', outputSuffix: 'ABSREL.json', @@ -754,23 +776,23 @@ // Extract meaningful error from HyPhy stdout for better error reporting let hyphyError = null; if (hyphyOut) { - const lines = hyphyOut.split('\n').filter(l => l.trim()); + const lines = hyphyOut.split('\n').filter((l) => l.trim()); const errorPatterns = [ /Expected \d+ sites?, but found \d+/, /stop codons? found/i, /not a valid/i, /divisible by 3/i, /at least 3 unique sequences/i, - /too large/i, + /too large/i ]; for (const line of lines) { - if (errorPatterns.some(p => p.test(line))) { + if (errorPatterns.some((p) => p.test(line))) { hyphyError = line.trim(); break; } } if (!hyphyError) { - const errorLine = lines.find(l => /^Error:/i.test(l.trim())); + const errorLine = lines.find((l) => /^Error:/i.test(l.trim())); if (errorLine) { hyphyError = errorLine.trim().replace(/^Error:\s*/i, ''); } @@ -778,7 +800,8 @@ } // Check for tree-related messages in datareader output (not an error, just informational) - const hasTreeParsingMessage = hyphyOut.includes('Illegal right hand side in call to Topology') || + const hasTreeParsingMessage = + hyphyOut.includes('Illegal right hand side in call to Topology') || hyphyOut.includes('tree string is invalid') || hyphyOut.includes('Newick tree spec'); @@ -788,9 +811,9 @@ // Log any error-like messages for debugging (but don't fail) if (hyphyOut.includes('Error') || hyphyOut.includes('error')) { - const errorLines = hyphyOut.split('\n').filter(line => - line.toLowerCase().includes('error') - ); + const errorLines = hyphyOut + .split('\n') + .filter((line) => line.toLowerCase().includes('error')); if (errorLines.length > 0) { console.warn('HyPhy datareader messages:', errorLines); } @@ -810,7 +833,10 @@ jsonBlob = await cliObj.download('/shared/data/results.json'); } catch (downloadError) { console.error('Failed to download results.json:', downloadError); - throw new Error(hyphyError || 'File analysis failed. The file may be in an unsupported format or contain invalid data.'); + throw new Error( + hyphyError || + 'File analysis failed. The file may be in an unsupported format or contain invalid data.' + ); } const response = await fetch(jsonBlob); const blob = await response.blob(); @@ -819,15 +845,26 @@ // Validate that we got JSON, not an error page if (jsonOut.trim().startsWith('= 0 ? 'codon' : fileMetricsJSON.FILE_INFO?.gencodeid === -1 ? 'nucleotide' : 'protein', + format: + fileMetricsJSON.FILE_INFO?.gencodeid >= 0 + ? 'codon' + : fileMetricsJSON.FILE_INFO?.gencodeid === -1 + ? 'nucleotide' + : 'protein', sequenceCount: fileMetricsJSON.FILE_INFO?.sequences || 0, siteCount: fileMetricsJSON.FILE_INFO?.sites || 0 }); @@ -936,15 +978,25 @@ const errorMsg = error.message || ''; const isTreeError = /Tree|tree|Topology|Newick/.test(errorMsg); - const errorType = isTreeError ? 'invalid-tree' - : errorMsg.includes('divisible by 3') ? 'not-codon-aligned' - : errorMsg.includes('stop codon') ? 'stop-codons' - : /nucleotide alignment|protein alignment|character states/.test(errorMsg) ? 'wrong-alphabet' - : /too large|must include prebuilt/.test(errorMsg) ? 'dataset-too-large' - : /at least \d+ unique sequences|No sequences found|No MATRIX block/.test(errorMsg) ? 'insufficient-sequences' - : /Aioli|not a valid value for parameter|Invalid parameter choice/.test(errorMsg) ? 'runtime-error' - : /format|valid alignment|valid FASTA|valid NEXUS|valid sequence|FASTA data is empty|partition specification|Sequence data found before header|File is empty/i.test(errorMsg) ? 'invalid-format' - : 'unknown'; + const errorType = isTreeError + ? 'invalid-tree' + : errorMsg.includes('divisible by 3') + ? 'not-codon-aligned' + : errorMsg.includes('stop codon') + ? 'stop-codons' + : /nucleotide alignment|protein alignment|character states/.test(errorMsg) + ? 'wrong-alphabet' + : /too large|must include prebuilt/.test(errorMsg) + ? 'dataset-too-large' + : /at least \d+ unique sequences|No sequences found|No MATRIX block/.test(errorMsg) + ? 'insufficient-sequences' + : /Aioli|not a valid value for parameter|Invalid parameter choice/.test(errorMsg) + ? 'runtime-error' + : /format|valid alignment|valid FASTA|valid NEXUS|valid sequence|FASTA data is empty|partition specification|Sequence data found before header|File is empty/i.test( + errorMsg + ) + ? 'invalid-format' + : 'unknown'; const eventPayload = { errorType, stage: currentStage, source }; if (errorType === 'unknown' || errorType === 'invalid-format') { // Defensive fallback: 338/342 unknown events had no message field in @@ -1145,7 +1197,7 @@
-

- Introducing PRIME -

-

- Property-Informed Models of Evolution -

+

Introducing PRIME

+

Property-Informed Models of Evolution

- Characterize physicochemical selection in protein evolution. PRIME incorporates amino acid properties like molecular volume, hydropathy, and structural propensities to resolve the biophysical basis of selective constraint. + Characterize physicochemical selection in protein evolution. PRIME incorporates + amino acid properties like molecular volume, hydropathy, and structural propensities + to resolve the biophysical basis of selective constraint.

+
{/if} - + {#if activeTab === 'data'} diff --git a/src/test/axomeme-assemble.test.js b/src/test/axomeme-assemble.test.js new file mode 100644 index 0000000..3e05d29 --- /dev/null +++ b/src/test/axomeme-assemble.test.js @@ -0,0 +1,269 @@ +/** + * Tests for AxoMEME tensor assembly. + * + * Assembly is where the separately-verified stages get joined, so the failures available here are + * joining failures: right values in the wrong order, right shapes with the wrong species, MDS + * computed on the wrong matrix. None of them crash. All of them produce a well-formed bundle that + * means something the model was not trained on, which is why every bundle below is also run through + * validateInputBundle. + */ +import { describe, it, expect } from 'vitest'; +import { + prepareAlignment, + chooseReference, + orderSpecies, + batchSizeFor +} from '../lib/services/axomeme/assemble.js'; +import { parseNewick } from '../lib/services/axomeme/newick.js'; +import { computeMdsCoordinates } from '../lib/services/axomeme/mds.js'; +import { codonToken, aaToken } from '../lib/services/axomeme/tokenizer.js'; +import { + validateInputBundle, + CODON_UNKNOWN, + AA_UNKNOWN, + MDS_COMPONENTS +} from '../lib/services/axomeme/modelContract.js'; + +/** Four taxa, distinct codons, 3 sites each. Tree order is deliberately NOT alignment order. */ +const NAMES = ['alpha', 'beta', 'gamma', 'delta']; +const SEQS = ['ATGTTATCA', 'ATGCTATCA', 'ATGTTAAGC', 'ATGGGGTCA']; +const TREE = '((gamma:0.1,delta:0.2):0.05,(beta:0.3,alpha:0.15):0.02);'; + +const prep = (over = {}) => + prepareAlignment({ names: NAMES, sequences: SEQS, treeText: TREE, maxSpecies: 8, ...over }); + +describe('chooseReference', () => { + it('honours an explicit choice', () => { + expect(chooseReference(NAMES, 'gamma')).toBe('gamma'); + }); + + it('falls back to the first sequence, which is what fires on viral data', () => { + // The heuristic looks for 'hg' / 'hg38' / 'human' — a TOGA-mammal artifact. DataMonkey traffic + // is viral, so the fallback is the real behaviour. + expect(chooseReference(NAMES)).toBe('alpha'); + expect(chooseReference(['x', 'human', 'y'])).toBe('human'); + expect(chooseReference(['x', 'hg38'])).toBe('hg38'); + }); + + it('ignores an explicit name that is not in the alignment', () => { + expect(chooseReference(NAMES, 'nope')).toBe('alpha'); + }); +}); + +describe('orderSpecies', () => { + it('takes TREE order, not alignment order', () => { + // Order fixes the distance matrix rows and therefore MDS, and index 0 seeds Max-PD. + const tree = parseNewick(TREE); + const { order, matchedFromTree } = orderSpecies(NAMES, tree, 'gamma'); + expect(matchedFromTree).toBe(true); + expect(order.map((i) => NAMES[i])).toEqual(['gamma', 'delta', 'beta', 'alpha']); + }); + + it('moves the reference sequence to the front', () => { + const tree = parseNewick(TREE); + const { order } = orderSpecies(NAMES, tree, 'alpha'); + expect(order.map((i) => NAMES[i])[0]).toBe('alpha'); + // and everything else keeps tree order behind it + expect(order.map((i) => NAMES[i])).toEqual(['alpha', 'gamma', 'delta', 'beta']); + }); + + it('falls back to alignment order when nothing matches the tree', () => { + const tree = parseNewick('((zzz:0.1,yyy:0.2):0.05);'); + const { order, matchedFromTree } = orderSpecies(NAMES, tree, 'alpha'); + expect(matchedFromTree).toBe(false); + expect(order.map((i) => NAMES[i])).toEqual(NAMES); + }); + + it('drops alignment sequences that are absent from the tree', () => { + const tree = parseNewick('((gamma:0.1,delta:0.2):0.05);'); + const { order } = orderSpecies(NAMES, tree, 'gamma'); + expect(order.map((i) => NAMES[i])).toEqual(['gamma', 'delta']); + }); +}); + +describe('prepareAlignment', () => { + it('produces a bundle that satisfies the contract', () => { + const p = prep(); + const bundle = p.batch(0); + const v = validateInputBundle(bundle, { + batch: p.totalCodons, + numSpecies: p.speciesCount, + windowSize: p.windowSize + }); + expect(v.errors).toEqual([]); + }); + + it('derives the site count from the reference sequence', () => { + expect(prep().totalCodons).toBe(3); // 9 nt / 3 + }); + + it('feeds the graph only the real species, with nothing padded', () => { + // Measured equivalent to feeding max_species with the remainder masked (1 float32 ulp), and + // the difference is 462 MB vs 2.3 MB on a 441-site alignment. + const p = prep(); + expect(p.speciesCount).toBe(4); + expect(Array.from(p.paddingMask)).toEqual([0, 0, 0, 0]); + expect(p.batch(0).dist_matrix.dims).toEqual([3, 4, 4]); + }); + + it('computes MDS on the PADDED matrix and slices, not on the real N', () => { + // The single most silently-wrong thing available in this file. Coordinates depend on + // max_species because the padded zeros take part in the double-centring. + const p = prepareAlignment({ + names: NAMES, + sequences: SEQS, + treeText: TREE, + maxSpecies: 16 + }); + const cap = 16; + const padded = new Float64Array(cap * cap); + const n = p.speciesCount; + for (let i = 0; i < n; i++) { + for (let j = 0; j < n; j++) padded[i * cap + j] = p.dist[i * n + j]; + } + const expected = computeMdsCoordinates(padded, cap, MDS_COMPONENTS); + for (let i = 0; i < n * MDS_COMPONENTS; i++) { + expect(p.mds[i]).toBeCloseTo(expected[i], 6); + } + // And it is genuinely different from the unpadded answer, so the test above has teeth. + const unpadded = computeMdsCoordinates(Float64Array.from(p.dist), n, MDS_COMPONENTS); + expect(Array.from(p.mds)).not.toEqual(Array.from(unpadded)); + }); + + it('orders the distance matrix rows to match the selected species', () => { + const p = prep({ referenceName: 'gamma' }); + expect(p.selectedNames).toEqual(['gamma', 'delta', 'beta', 'alpha']); + const n = p.speciesCount; + // gamma-delta share a parent: 0.1 + 0.2 = 0.3. gamma-beta crosses the root: 0.1+0.05+0.02+0.3. + expect(p.dist[0 * n + 1]).toBeCloseTo(0.3, 5); + expect(p.dist[0 * n + 2]).toBeCloseTo(0.47, 5); + expect(p.dist[0 * n + 0]).toBe(0); + }); + + it('tokenises each species at each site, in selected order', () => { + const p = prep({ referenceName: 'alpha' }); + // selected: alpha, gamma, delta, beta -> site 1 (2nd codon) is TTA, TTA, GGG, CTA + const n = p.speciesCount; + const at = (site, s) => Number(p.codonTokens[(site * n + s) * p.windowSize]); + expect(at(1, 0)).toBe(codonToken('TTA')); + expect(at(1, 1)).toBe(codonToken('TTA')); + expect(at(1, 2)).toBe(codonToken('GGG')); + expect(at(1, 3)).toBe(codonToken('CTA')); + const aaAt = (site, s) => Number(p.aaTokens[(site * n + s) * p.windowSize]); + expect(aaAt(1, 2)).toBe(aaToken('GGG')); + }); + + it('leaves pad values where a sequence is shorter than the reference', () => { + // `torch.ones(...) * 65` is never overwritten past a short sequence's end. + const p = prepareAlignment({ + names: ['a', 'b'], + sequences: ['ATGTTATCA', 'ATG'], + treeText: '(a:0.1,b:0.2);', + maxSpecies: 8 + }); + const n = p.speciesCount; + const idx = (site, s) => (site * n + s) * p.windowSize; + expect(Number(p.codonTokens[idx(0, 1)])).toBe(codonToken('ATG')); + expect(Number(p.codonTokens[idx(1, 1)])).toBe(CODON_UNKNOWN); + expect(Number(p.aaTokens[idx(1, 1)])).toBe(AA_UNKNOWN); + }); + + it('applies Max-PD when over the cap, seeded at the reference', () => { + const names = ['r', 'near', 'far', 'mid']; + const seqs = ['ATG', 'ATG', 'ATG', 'ATG']; + const p = prepareAlignment({ + names, + sequences: seqs, + treeText: '((r:0.01,near:0.01):0.05,(far:2.0,mid:0.5):0.5);', + maxSpecies: 2, + referenceName: 'r' + }); + expect(p.speciesCount).toBe(2); + expect(p.selectedNames[0]).toBe('r'); // the seed + expect(p.selectedNames[1]).toBe('far'); // farthest from it + }); + + it('clamps a negative distance to zero and REPORTS the magnitude', () => { + // DM3's own NJ emits negative branch lengths; the large.nex demo produces a patristic sum of + // -1.04e-5, which is zero with rounding error on it. The model was trained on clamped + // distances (the handoff README says so), so clamping matches training — but doing it silently + // would hide a genuinely broken tree, which is why the magnitude comes back out. + const p = prepareAlignment({ + names: ['a', 'b', 'c'], + sequences: ['ATG', 'ATG', 'ATG'], + treeText: '((a:-0.5,b:0.2):0.05,c:0.3);', + maxSpecies: 8 + }); + expect(Array.from(p.dist).every((v) => v >= 0)).toBe(true); + expect(p.clampedDistances).toBeGreaterThan(0); + expect(p.mostNegativeDistance).toBeCloseTo(-0.3, 5); // a-b: -0.5 + 0.2 + }); + + it('reports nothing clamped for a clean tree', () => { + const p = prep(); + expect(p.clampedDistances).toBe(0); + expect(p.mostNegativeDistance).toBe(0); + }); + + it('produces a contract-valid bundle from a tree with negative branch lengths', () => { + // The large.nex regression: the bundle used to be REJECTED by validateInputBundle because a + // patristic sum came out at -1e-5. Clamping is what makes an ordinary NJ tree usable. + const p = prepareAlignment({ + names: ['a', 'b', 'c', 'd'], + sequences: ['ATGTTA', 'ATGCTA', 'ATGGGG', 'ATGAAA'], + treeText: '((a:-0.00001,b:0.2):0.05,(c:0.3,d:0.1):0.02);', + maxSpecies: 8 + }); + const v = validateInputBundle(p.batch(0), { + batch: p.totalCodons, + numSpecies: p.speciesCount, + windowSize: p.windowSize + }); + expect(v.errors).toEqual([]); + }); + + it('falls back to an all-zero distance matrix without a tree', () => { + const p = prepareAlignment({ names: NAMES, sequences: SEQS, maxSpecies: 8 }); + expect(Array.from(p.dist).every((v) => v === 0)).toBe(true); + expect(p.matchedFromTree).toBe(false); + // MDS of an all-zero matrix is all zeros, not NaN. + expect(Array.from(p.mds).every((v) => v === 0)).toBe(true); + }); + + it('rejects mismatched or empty input rather than producing a bundle', () => { + expect(() => prepareAlignment({ names: ['a'], sequences: [] })).toThrow(/parallel/); + expect(() => prepareAlignment({ names: [], sequences: [] })).toThrow(/no sequences/); + }); +}); + +describe('batching', () => { + it('slices sites without disturbing the per-alignment tensors', () => { + const p = prep(); + const all = p.batch(0); + const tail = p.batch(1, 2); + expect(tail.msa_codons.dims).toEqual([2, 4, 1]); + // site 1 of the full batch is site 0 of this one + const n = p.speciesCount; + for (let s = 0; s < n; s++) { + expect(tail.msa_codons.data[s]).toBe(all.msa_codons.data[n + s]); + } + // and the invariant tensors are repeated per site, identically + for (let k = 0; k < n * n; k++) { + expect(tail.dist_matrix.data[n * n + k]).toBe(tail.dist_matrix.data[k]); + } + }); + + it('clamps a range that runs past the end', () => { + const p = prep(); + expect(p.batch(2, 99).msa_codons.dims[0]).toBe(1); + expect(p.batch(3, 5).msa_codons.dims[0]).toBe(0); + }); + + it('sizes batches against the dist_matrix budget', () => { + // 4 * N^2 bytes per site is the dominant term. + expect(batchSizeFor(512, 64 * 1024 * 1024)).toBe(64); + expect(batchSizeFor(36, 64 * 1024 * 1024)).toBeGreaterThan(1000); + // Never zero: one site of a huge alignment still has to go through. + expect(batchSizeFor(4096, 1024)).toBe(1); + }); +}); diff --git a/src/test/axomeme-mds.test.js b/src/test/axomeme-mds.test.js new file mode 100644 index 0000000..e42af1e --- /dev/null +++ b/src/test/axomeme-mds.test.js @@ -0,0 +1,221 @@ +/** + * Tests for the symmetric eigendecomposition and classical MDS. + * + * These check PROPERTIES, not recorded outputs. An eigendecomposition test that asserts "these are + * the numbers we got last time" proves determinism and nothing else; the real contract is + * A·V = V·Λ with V orthonormal, and that is checkable from first principles on any input. The MDS + * cases are hand-derived from small distance matrices whose answer can be worked out on paper. + * + * The cross-implementation check against numpy lives in scripts/axomeme/verify_preprocessing.py + + * .mjs, which compares against the ML team's own `compute_mds_coordinates` over 270 real DataMonkey + * trees. That run is what proves parity; these tests are what make a failure there interpretable. + */ +import { describe, it, expect } from 'vitest'; +import { symmetricEigen } from '../lib/services/axomeme/symmetricEigen.js'; +import { computeMdsCoordinates } from '../lib/services/axomeme/mds.js'; + +/** Deterministic symmetric matrix, fixed seed so a failure is reproducible. */ +function symMatrix(n, seed) { + let s = seed; + const rnd = () => ((s = (s * 1103515245 + 12345) & 0x7fffffff) / 0x7fffffff) * 2 - 1; + const A = new Float64Array(n * n); + for (let i = 0; i < n; i++) { + for (let j = i; j < n; j++) { + const v = rnd(); + A[i * n + j] = v; + A[j * n + i] = v; + } + } + return A; +} + +describe('symmetricEigen', () => { + it('solves a 2x2 with known eigenvalues', () => { + // [[2,1],[1,2]] has eigenvalues 1 and 3. + const { values } = symmetricEigen([2, 1, 1, 2], 2); + expect(values[0]).toBeCloseTo(1, 12); + expect(values[1]).toBeCloseTo(3, 12); + }); + + it('returns eigenvalues ASCENDING, matching numpy.linalg.eigh', () => { + // mds.js reads components from the END of this array; a descending convention would silently + // return the four SMALLEST components. + const { values } = symmetricEigen(symMatrix(20, 3), 20); + for (let i = 1; i < 20; i++) expect(values[i]).toBeGreaterThanOrEqual(values[i - 1]); + }); + + it('reconstructs the matrix: A = V Λ Vᵀ', () => { + for (const n of [3, 10, 40]) { + const A = symMatrix(n, n * 7 + 1); + const { values, vectors } = symmetricEigen(A, n); + let worst = 0; + for (let i = 0; i < n; i++) { + for (let j = 0; j < n; j++) { + let acc = 0; + for (let k = 0; k < n; k++) acc += vectors[i * n + k] * values[k] * vectors[j * n + k]; + worst = Math.max(worst, Math.abs(acc - A[i * n + j])); + } + } + expect(worst, `n=${n}`).toBeLessThan(1e-12); + } + }); + + it('produces orthonormal eigenvectors', () => { + const n = 30; + const { vectors } = symmetricEigen(symMatrix(n, 11), n); + let worst = 0; + for (let i = 0; i < n; i++) { + for (let j = 0; j < n; j++) { + let dot = 0; + for (let k = 0; k < n; k++) dot += vectors[k * n + i] * vectors[k * n + j]; + worst = Math.max(worst, Math.abs(dot - (i === j ? 1 : 0))); + } + } + expect(worst).toBeLessThan(1e-12); + }); + + it('handles a diagonal matrix, the identity, and n=1', () => { + const { values } = symmetricEigen([3, 0, 0, 0, 1, 0, 0, 0, 2], 3); + expect(Array.from(values)).toEqual([1, 2, 3]); // sorted ascending + const id = symmetricEigen([1, 0, 0, 1], 2); + expect(Array.from(id.values)).toEqual([1, 1]); + const one = symmetricEigen([7], 1); + expect(one.values[0]).toBe(7); + }); + + it('does not mutate the caller"s matrix', () => { + const A = Float64Array.from([2, 1, 1, 2]); + symmetricEigen(A, 2); + expect(Array.from(A)).toEqual([2, 1, 1, 2]); + }); + + it('survives a matrix with repeated eigenvalues', () => { + // 2I has a fully degenerate spectrum — any orthonormal basis is correct. The routine must + // still return orthonormal vectors and the right values rather than dividing by a zero gap. + const { values, vectors } = symmetricEigen([2, 0, 0, 0, 2, 0, 0, 0, 2], 3); + expect(Array.from(values)).toEqual([2, 2, 2]); + for (let i = 0; i < 3; i++) { + let norm = 0; + for (let k = 0; k < 3; k++) norm += vectors[k * 3 + i] ** 2; + expect(norm).toBeCloseTo(1, 12); + } + }); +}); + +describe('computeMdsCoordinates', () => { + it('returns zeros when n <= nComponents', () => { + // `if N <= n_components: return zeros` — note <=, so n === nComponents is also all zeros. + expect(Array.from(computeMdsCoordinates(new Float64Array(9), 3, 4))).toEqual( + new Array(12).fill(0) + ); + expect(Array.from(computeMdsCoordinates(new Float64Array(16), 4, 4))).toEqual( + new Array(16).fill(0) + ); + }); + + it('recovers collinear points from their distances', () => { + // Three points at 0, 1, 2 on a line. Worked by hand: D2 double-centres to + // [[1,0,-1],[0,0,0],[-1,0,1]], whose only positive eigenvalue is 2 with eigenvector + // [1,0,-1]/sqrt(2), so component 0 is [1, 0, -1] and component 1 is all zeros. + const D = [0, 1, 2, 1, 0, 1, 2, 1, 0]; + const c = computeMdsCoordinates(D, 3, 2); + expect(c[0 * 2 + 0]).toBeCloseTo(1, 5); + expect(c[1 * 2 + 0]).toBeCloseTo(0, 5); + expect(c[2 * 2 + 0]).toBeCloseTo(-1, 5); + // The second component's eigenvalue is mathematically ZERO — three collinear points need one + // dimension. It does not come out as exactly zero, though: it lands on float dust around + // 1e-17, and the reference's guard is `if val > 0`, which does not distinguish a true zero + // from dust. So a coordinate of ~1e-8 (sqrt of the dust) is emitted rather than a clean zero. + // That is the reference's behaviour and this port matches it; asserting `toBe(0)` here would + // be asserting something neither implementation does. + for (const i of [0, 1, 2]) expect(Math.abs(c[i * 2 + 1])).toBeLessThan(1e-6); + }); + + it('preserves pairwise distances for points that embed exactly', () => { + // A square of side 1: MDS in 2 dimensions must reproduce the input distances. + const s2 = Math.SQRT2; + const D = [0, 1, s2, 1, 1, 0, 1, s2, s2, 1, 0, 1, 1, s2, 1, 0]; + const c = computeMdsCoordinates(D, 4, 2); + const dist = (i, j) => Math.hypot(c[i * 2] - c[j * 2], c[i * 2 + 1] - c[j * 2 + 1]); + expect(dist(0, 1)).toBeCloseTo(1, 4); + expect(dist(1, 2)).toBeCloseTo(1, 4); + expect(dist(0, 2)).toBeCloseTo(s2, 4); + expect(dist(1, 3)).toBeCloseTo(s2, 4); + }); + + it('applies the sign convention: the largest-magnitude entry is positive', () => { + // The reference's rule, and the thing that removes half the eigenvector ambiguity for free. + const D = [0, 1, 2, 1, 0, 1, 2, 1, 0]; + const c = computeMdsCoordinates(D, 3, 2); + let maxAbs = 0; + let atIdx = 0; + for (let i = 0; i < 3; i++) { + if (Math.abs(c[i * 2]) > maxAbs) { + maxAbs = Math.abs(c[i * 2]); + atIdx = i; + } + } + expect(c[atIdx * 2]).toBeGreaterThan(0); + }); + + it('DEPENDS ON THE PADDING, because the reference runs MDS on the padded matrix', () => { + // Not a quirk to be optimised away. The padded zeros take part in the double-centring, so the + // same real taxa padded to different max_species produce different coordinates. A port that + // runs MDS on the real N and pads afterwards silently feeds the model different inputs. + const real = [0, 1, 2, 1, 0, 1, 2, 1, 0]; + const at = (cap) => { + const p = new Float64Array(cap * cap); + for (let i = 0; i < 3; i++) for (let j = 0; j < 3; j++) p[i * cap + j] = real[i * 3 + j]; + return computeMdsCoordinates(p, cap, 2); + }; + const a = at(8); + const b = at(16); + // Component 1 is where it shows. Measured across caps for these three collinear points: + // cap= 3 -> 9.50e-9 cap= 4 -> -9.833e-2 cap= 8 -> -1.0097e-1 + // cap=16 -> -1.0133e-1 cap=64 -> -1.0150e-1 + // It converges as the padding grows but never stops depending on it. + expect(a[0 * 2 + 1]).not.toBeCloseTo(b[0 * 2 + 1], 6); + expect(a[1 * 2 + 1]).not.toBeCloseTo(b[1 * 2 + 1], 6); + // Component 0's MAGNITUDE happens to be padding-invariant for this symmetric example (the + // collinear geometry dominates), which is why the check above is on component 1 — picking + // component 0 would have made this test pass for the wrong reason and then fail later. + expect(Math.abs(a[0])).toBeCloseTo(Math.abs(b[0]), 6); + }); + + it('rounds distances to float32 before squaring, as the reference does', () => { + // The reference is handed a float32 torch tensor, so it squares float32 values. Feeding + // float64 changes components 2-3 by up to 99% on real trees (measured), because squared + // distances reach ~1e6 while the fourth eigenvalue can be ~1e-1. Passing an already-rounded + // matrix and a full-precision one must therefore give the SAME answer. + const n = 8; + const raw = new Float64Array(n * n); + let s = 3; + const rnd = () => (s = (s * 1103515245 + 12345) & 0x7fffffff) / 0x7fffffff; + for (let i = 0; i < n; i++) { + for (let j = i + 1; j < n; j++) { + const v = rnd() * 1000; + raw[i * n + j] = v; + raw[j * n + i] = v; + } + } + const rounded = Float64Array.from(raw, (v) => Math.fround(v)); + const a = computeMdsCoordinates(raw, n, 4); + const b = computeMdsCoordinates(rounded, n, 4); + expect(Array.from(a)).toEqual(Array.from(b)); + }); + + it('returns float32 values, matching the reference cast', () => { + const D = [0, 1, 2, 1, 0, 1, 2, 1, 0]; + const c = computeMdsCoordinates(D, 3, 2); + expect(c).toBeInstanceOf(Float32Array); + for (const v of c) expect(Math.fround(v)).toBe(v); + }); + + it('treats a negative distance as its magnitude, because squaring loses the sign', () => { + // Documented consequence rather than desired behaviour: DM3's NJ emits negative branch + // lengths, and MDS silently absorbs them where the reference's density term throws. + const pos = computeMdsCoordinates([0, 1, 2, 1, 0, 1, 2, 1, 0], 3, 2); + const neg = computeMdsCoordinates([0, -1, 2, -1, 0, 1, 2, 1, 0], 3, 2); + expect(Array.from(neg)).toEqual(Array.from(pos)); + }); +}); diff --git a/src/test/axomeme-model-contract.test.js b/src/test/axomeme-model-contract.test.js new file mode 100644 index 0000000..e7f056e --- /dev/null +++ b/src/test/axomeme-model-contract.test.js @@ -0,0 +1,246 @@ +/** + * Tests for the AxoMEME 2.0 ONNX input contract. + * + * These are not shape-checking-the-shape-checker busywork. Every case below is a mistake a JS + * preprocessing port actually makes — one that produces a tensor of exactly the right dtype and + * exactly the right dimensions, and means something the model was never trained on. Those are the + * errors that cost days, because nothing crashes: the graph runs, five numbers come out per site, + * and they are wrong in a way that looks like a bad model rather than a bad tensor. + * + * The constants themselves are pinned too. They are transcribed from the ML team's handoff scripts, + * and a transcription error in, say, the codon vocabulary order is invisible until it is compared + * against real fixtures — which is a much later and much more expensive place to find it. + */ +import { describe, it, expect } from 'vitest'; +import { + CODON_ORDER, + CODON_GAP, + CODON_UNKNOWN, + NUM_CODON_TOKENS, + AA_LIST, + AA_GAP, + AA_UNKNOWN, + CODON_VALID_BELOW, + AA_VALID_BELOW, + MAX_SPECIES_DEFAULT, + WINDOW_SIZE_DEFAULT, + MDS_COMPONENTS, + INPUT_SPEC, + INPUT_NAMES, + OUTPUT_SPEC, + VERIFIED_MODEL_SHA256, + validateInputBundle +} from '../lib/services/axomeme/modelContract.js'; + +const BATCH = 2; +const SPECIES = 4; // index 3 is padded +const WIN = 1; + +/** A bundle that satisfies the contract. Each test breaks exactly one thing about it. */ +function validBundle() { + const d = [ + [0, 0.1, 0.2, 0], + [0.1, 0, 0.3, 0], + [0.2, 0.3, 0, 0], + [0, 0, 0, 0] + ]; + const distOne = d.flat(); + return { + msa_codons: { + data: new BigInt64Array([0n, 5n, 63n, 65n, 1n, 2n, 3n, 65n]), + dims: [BATCH, SPECIES, WIN] + }, + msa_aas: { + data: new BigInt64Array([0n, 4n, 20n, 22n, 1n, 2n, 3n, 22n]), + dims: [BATCH, SPECIES, WIN] + }, + dist_matrix: { + data: new Float32Array([...distOne, ...distOne]), + dims: [BATCH, SPECIES, SPECIES] + }, + mds_coords: { + data: new Float32Array(BATCH * SPECIES * MDS_COMPONENTS), + dims: [BATCH, SPECIES, MDS_COMPONENTS] + }, + // TRUE = padded. Species 3 only. + padding_mask: { + data: new Uint8Array([0, 0, 0, 1, 0, 0, 0, 1]), + dims: [BATCH, SPECIES] + } + }; +} + +const check = (b) => validateInputBundle(b, { batch: BATCH, numSpecies: SPECIES, windowSize: WIN }); + +describe('the transcribed constants', () => { + it('orders codons TCAG, not alphabetically', () => { + // The single most damaging transcription error available here: an ACGT vocabulary is a valid + // permutation of the same 64 tokens and is wrong at every site of every alignment. + expect(CODON_ORDER).toBe('TCAG'); + const codons = [...CODON_ORDER].flatMap((a) => + [...CODON_ORDER].flatMap((b) => [...CODON_ORDER].map((c) => a + b + c)) + ); + expect(codons).toHaveLength(64); + expect(codons[0]).toBe('TTT'); + expect(codons[63]).toBe('GGG'); + // If someone "fixes" the order to ACGT this is the assertion that objects. + // Met. Under TCAG: A=2, T=0, G=3 -> 2*16 + 0*4 + 3 = 35. Under ACGT it would be 14, so this + // single number distinguishes the two orderings. + expect(codons.indexOf('ATG')).toBe(35); + }); + + it('keeps gap and unknown distinct, and different between the two streams', () => { + expect(CODON_GAP).toBe(64); + expect(CODON_UNKNOWN).toBe(65); + expect(NUM_CODON_TOKENS).toBe(66); + expect(AA_GAP).toBe(AA_LIST.indexOf('-')); + expect(AA_UNKNOWN).toBe(AA_LIST.indexOf('?')); + expect(AA_GAP).toBe(21); + expect(AA_UNKNOWN).toBe(22); + expect(CODON_GAP).not.toBe(AA_GAP); + }); + + it('sets the validity thresholds so that a GAP is not a valid observation', () => { + // forward() gates on (c < 64) & (a < 21). A gap is 64 / 21, so it fails both — deliberately. + expect(CODON_VALID_BELOW).toBe(CODON_GAP); + expect(AA_VALID_BELOW).toBe(AA_GAP); + expect(CODON_GAP < CODON_VALID_BELOW).toBe(false); + expect(AA_GAP < AA_VALID_BELOW).toBe(false); + }); + + it('pins the checkpoint defaults', () => { + expect(MAX_SPECIES_DEFAULT).toBe(512); + expect(MDS_COMPONENTS).toBe(4); + // window_size 1 means the central index is 0 and every window IS the site. An even window + // would put the scored codon off-centre. + expect(WINDOW_SIZE_DEFAULT).toBe(1); + expect(Math.floor(WINDOW_SIZE_DEFAULT / 2)).toBe(0); + }); + + it('lists the five inputs in forward() order and marks the site-invariant ones', () => { + expect(INPUT_NAMES).toEqual([ + 'msa_codons', + 'msa_aas', + 'dist_matrix', + 'mds_coords', + 'padding_mask' + ]); + // These three are computed once per alignment and expanded across sites. That is what makes + // batching every site into a single graph run cheap, so it is worth asserting rather than + // rediscovering. + const invariant = INPUT_SPEC.filter((s) => s.siteInvariant).map((s) => s.name); + expect(invariant).toEqual(['dist_matrix', 'mds_coords', 'padding_mask']); + }); + + it('names the five outputs the graph actually exposes', () => { + // Read from InferenceSession.outputNames on the real artifact, not from its README. This was an + // open question — the shipped driver reaches for the TRAIN branch to get raw ordinal logits — + // and the answer is that the export took the EVAL branch, so `lrt` arrives already decoded. + expect(OUTPUT_SPEC.map((o) => o.name)).toEqual([ + 'lrt', + 'alpha', + 'beta_neg', + 'beta_pos', + 'p_neg' + ]); + expect(OUTPUT_SPEC[0].note).toMatch(/already applied in-graph/); + }); + + it('records that the rate heads are log1p and need expm1 downstream', () => { + // The heads are softplus, and expm1(softplus(x)) == exp(x) — so these outputs are log1p(rate), + // not rates. Rendering them raw would understate every rate. p_neg is the exception. + for (const name of ['alpha', 'beta_neg', 'beta_pos']) { + expect(OUTPUT_SPEC.find((o) => o.name === name).note, name).toMatch(/expm1/); + } + expect(OUTPUT_SPEC.find((o) => o.name === 'p_neg').note).toMatch(/sigmoid/); + }); + + it('pins the artifact the contract was verified against', () => { + // A different export is not necessarily wrong, but the eval-mode conclusion was read off THIS + // graph, so swapping the model without revisiting this file is a mistake worth failing on. + expect(VERIFIED_MODEL_SHA256).toMatch(/^[0-9a-f]{64}$/); + }); + + it('cannot be mutated by a caller', () => { + expect(() => { + INPUT_SPEC.push({ name: 'nope' }); + }).toThrow(); + }); +}); + +describe('validateInputBundle', () => { + it('accepts a well-formed bundle', () => { + const r = check(validBundle()); + expect(r.errors).toEqual([]); + expect(r.ok).toBe(true); + }); + + it('reports a missing tensor by name', () => { + const b = validBundle(); + delete b.mds_coords; + expect(check(b).errors).toContain('mds_coords: missing'); + }); + + it('catches a transposed distance matrix shape', () => { + const b = validBundle(); + b.dist_matrix.dims = [BATCH, SPECIES, MDS_COMPONENTS + 1]; + expect(check(b).ok).toBe(false); + }); + + it('catches dims that are right but data that is short', () => { + // The shape says one thing and the buffer says another — onnxruntime will happily read past + // the end of the meaningful data or throw something opaque. + const b = validBundle(); + b.mds_coords.data = new Float32Array(4); + expect(check(b).errors.join(' ')).toMatch(/mds_coords: 4 elements/); + }); + + it('catches an out-of-vocabulary token', () => { + const b = validBundle(); + b.msa_codons.data = new BigInt64Array([0n, 5n, 63n, 66n, 1n, 2n, 3n, 65n]); + expect(check(b).errors.join(' ')).toMatch(/msa_codons\[3\] = 66/); + }); + + it('catches a FLIPPED padding mask — the error that passes every shape check', () => { + // This is the one worth having the whole validator for. `padding_mask` is TRUE for padded + // rows, which is the opposite of the "1 = keep" convention most attention APIs use. Invert it + // and every tensor is still perfectly well-formed; the model simply masks out every real + // taxon and attends to nothing. + const b = validBundle(); + b.padding_mask.data = new Uint8Array([1, 1, 1, 0, 1, 1, 1, 0]); + const r = check(b); + expect(r.ok).toBe(false); + expect(r.errors.join(' ')).toMatch(/TRUE = PADDED/); + }); + + it('catches a negative patristic distance, which is a real DM3 tree and a Python crash', () => { + // 5% of real DM3 trees carry a branch length <= -0.1; the Python inference path throws on + // them at predict_regression_nexus.py:955 rather than degrading. See treeSanitation.js. + const b = validBundle(); + b.dist_matrix.data[1] = -0.4; + const r = check(b); + expect(r.ok).toBe(false); + expect(r.errors.join(' ')).toMatch(/negative patristic distance/); + }); + + it('catches NaN before it reaches the graph', () => { + const b = validBundle(); + b.dist_matrix.data[2] = NaN; + expect(check(b).errors.join(' ')).toMatch(/dist_matrix\[2\] is NaN/); + }); + + it('catches a nonzero self-distance, which means the matrix is not a distance matrix', () => { + const b = validBundle(); + b.dist_matrix.data[0] = 0.5; // d(0,0) + expect(check(b).errors.join(' ')).toMatch(/self-distance/); + }); + + it('reports every independent problem, not just the first', () => { + // A port under development usually has several at once; stopping at the first costs a whole + // round trip per error. + const b = validBundle(); + delete b.msa_aas; + b.mds_coords.dims = [BATCH, SPECIES, 3]; + expect(check(b).errors.length).toBeGreaterThanOrEqual(2); + }); +}); diff --git a/src/test/axomeme-patristic.test.js b/src/test/axomeme-patristic.test.js new file mode 100644 index 0000000..ed5946e --- /dev/null +++ b/src/test/axomeme-patristic.test.js @@ -0,0 +1,265 @@ +/** + * Tests for the newick parser and patristic distances — the first half of the AxoMEME preprocessing + * port. + * + * The distances here are hand-computed from the newick, not recorded from this code. That + * distinction is the whole point: an expectation captured from the implementation asserts only that + * the implementation is deterministic, which it would be even if the arithmetic were wrong. Every + * number below can be checked by reading the tree string. + * + * The cross-implementation check against Python lives in scripts/axomeme/verify_preprocessing.py and + * is what actually proves parity; these tests are what make a failure there interpretable. + */ +import { describe, it, expect } from 'vitest'; +import { parseNewick, leafIndex, normalizeTaxonName } from '../lib/services/axomeme/newick.js'; +import { + rootDistances, + patristicRow, + patristicMatrix, + maxPdSelect +} from '../lib/services/axomeme/patristic.js'; + +/** ((A:0.1,B:0.2):0.05,C:0.3); — the worked example used throughout. */ +const SIMPLE = '((A:0.1,B:0.2):0.05,C:0.3);'; + +/** Node index of the leaf named `n`. */ +const leafOf = (tree, n) => leafIndex(tree).index.get(n); + +describe('parseNewick', () => { + it('builds the topology of a nested tree', () => { + const t = parseNewick(SIMPLE); + const { index } = leafIndex(t); + expect([...index.keys()].sort()).toEqual(['A', 'B', 'C']); + // A and B share a parent; C hangs off the root. + expect(t.parent[index.get('A')]).toBe(t.parent[index.get('B')]); + expect(t.parent[index.get('C')]).toBe(t.root); + expect(t.parent[t.parent[index.get('A')]]).toBe(t.root); + }); + + it('handles arbitrary nesting on both sides', () => { + const t = parseNewick('((A,B),(C,D));'); + const { index } = leafIndex(t); + expect([...index.keys()].sort()).toEqual(['A', 'B', 'C', 'D']); + expect(t.parent[index.get('A')]).toBe(t.parent[index.get('B')]); + expect(t.parent[index.get('C')]).toBe(t.parent[index.get('D')]); + expect(t.parent[index.get('A')]).not.toBe(t.parent[index.get('C')]); + }); + + it('nests a clade that appears after a leaf', () => { + // (A,(B,C)) exercises the branch where ',' and '(' arrive back to back. + const t = parseNewick('(A,(B,C));'); + const { index } = leafIndex(t); + expect(t.parent[index.get('A')]).toBe(t.root); + expect(t.parent[index.get('B')]).toBe(t.parent[index.get('C')]); + expect(t.parent[index.get('B')]).not.toBe(t.root); + }); + + it('reads plain, scientific and NEGATIVE branch lengths', () => { + const t = parseNewick('(a:1.5e-2,b:3.0E-3,c:-0.5);'); + const { index } = leafIndex(t); + expect(t.branchLength[index.get('a')]).toBeCloseTo(0.015, 12); + expect(t.branchLength[index.get('b')]).toBeCloseTo(0.003, 12); + // Preserved, NOT clamped — DM3's own NJ emits these and hiding them here would turn a loud + // downstream failure into a quiet wrong answer. + expect(t.branchLength[index.get('c')]).toBe(-0.5); + }); + + it('treats a missing branch length as 0, matching `branch_length or 0.0`', () => { + const t = parseNewick('((A,B),C);'); + expect(Array.from(t.branchLength).every((v) => v === 0)).toBe(true); + expect(Array.from(rootDistances(t)).every((v) => v === 0)).toBe(true); + }); + + it('does not mistake a bootstrap value for a taxon', () => { + // )95: is a label on an internal node. Reading it as a name invents a species. + const t = parseNewick('((A:0.1,B:0.2)95:0.05,C:0.3);'); + const { index } = leafIndex(t); + expect([...index.keys()].sort()).toEqual(['A', 'B', 'C']); + expect(index.has('95')).toBe(false); + // The label is kept on the internal node, just not treated as a leaf. + expect(t.name[t.parent[index.get('A')]]).toBe('95'); + }); + + it('keeps a quoted label containing a colon intact', () => { + // The entire reason newick quoting exists, and the case a naive split on ':' corrupts. + const t = parseNewick("(('Homo:sapiens':0.1,b:0.2):0.05);"); + const { index } = leafIndex(t); + expect(index.has('Homo:sapiens')).toBe(true); + expect(t.branchLength[index.get('Homo:sapiens')]).toBeCloseTo(0.1, 12); + }); + + it('strips newick comments rather than reading them as labels', () => { + const t = parseNewick('((A[&&NHX:x=1]:0.1,B:0.2):0.05,C:0.3);'); + expect([...leafIndex(t).index.keys()].sort()).toEqual(['A', 'B', 'C']); + }); + + it('parses without a trailing semicolon, and a bare single taxon', () => { + expect([...leafIndex(parseNewick('(A:0.1,B:0.2)')).index.keys()].sort()).toEqual(['A', 'B']); + const solo = parseNewick('A:0.1;'); + expect([...leafIndex(solo).index.keys()]).toEqual(['A']); + }); + + it('rejects empty input rather than returning an empty tree', () => { + expect(() => parseNewick('')).toThrow(/empty/); + expect(() => parseNewick(' ')).toThrow(/empty/); + }); + + it('lists leaves in PREORDER, matching Biopython get_terminals()', () => { + // Not cosmetic. The reference resolves duplicate tip names by dict overwrite, so the winner is + // the last leaf in THIS order; a breadth-first walk yields the same set and a different winner. + const t = parseNewick('((A:0.1,B:0.2):0.05,C:0.3);'); + expect(t.leaves.map((n) => t.name[n])).toEqual(['A', 'B', 'C']); + const t2 = parseNewick('(A,((B,C),D));'); + expect(t2.leaves.map((n) => t2.name[n])).toEqual(['A', 'B', 'C', 'D']); + }); + + it('resolves a duplicate tip name to the LAST leaf, and reports it', () => { + // Matches `{leaf.name: leaf for leaf in leaves}` — later entries overwrite earlier. Measured: + // 3 of 270 real DM3 trees have duplicate tips, and first-wins disagreed with the reference on + // every one of them. `duplicates` is what lets a caller refuse; the index itself stays + // faithful. + const t = parseNewick('((A:0.1,A:0.2):0.05,C:0.3);'); + const { index, duplicates } = leafIndex(t); + expect(duplicates).toEqual(['A']); + expect(index.size).toBe(2); + // The SECOND 'A' — the one with branch length 0.2. + expect(t.branchLength[index.get('A')]).toBeCloseTo(0.2, 12); + }); +}); + +describe('normalizeTaxonName', () => { + it('removes every quote anywhere, matching the reference', () => { + // Python's str.replace removes ALL occurrences; the reference does + // s.replace("'", "").replace('"', '').strip(). + expect(normalizeTaxonName("'Homo sapiens'")).toBe('Homo sapiens'); + expect(normalizeTaxonName("Homo_'sapiens'")).toBe('Homo_sapiens'); + expect(normalizeTaxonName(' "x" ')).toBe('x'); + }); +}); + +describe('rootDistances and patristic distances', () => { + it('accumulates root distances down the tree', () => { + const t = parseNewick(SIMPLE); + const d = rootDistances(t); + expect(d[leafOf(t, 'A')]).toBeCloseTo(0.15, 12); // 0.05 + 0.1 + expect(d[leafOf(t, 'B')]).toBeCloseTo(0.25, 12); // 0.05 + 0.2 + expect(d[leafOf(t, 'C')]).toBeCloseTo(0.3, 12); + expect(d[t.root]).toBe(0); + }); + + it('computes hand-checkable pairwise distances', () => { + const t = parseNewick(SIMPLE); + const nodes = ['A', 'B', 'C'].map((n) => leafOf(t, n)); + const m = patristicMatrix(t, nodes); + const at = (i, j) => m[i * 3 + j]; + expect(at(0, 1)).toBeCloseTo(0.3, 12); // A-B: 0.1 + 0.2 + expect(at(0, 2)).toBeCloseTo(0.45, 12); // A-C: 0.1 + 0.05 + 0.3 + expect(at(1, 2)).toBeCloseTo(0.55, 12); // B-C: 0.2 + 0.05 + 0.3 + }); + + it('is symmetric with a zero diagonal', () => { + const t = parseNewick('((A:0.1,B:0.2):0.05,(C:0.3,D:0.15):0.02);'); + const nodes = ['A', 'B', 'C', 'D'].map((n) => leafOf(t, n)); + const m = patristicMatrix(t, nodes); + for (let i = 0; i < 4; i++) { + expect(m[i * 4 + i]).toBe(0); + for (let j = 0; j < 4; j++) expect(m[i * 4 + j]).toBeCloseTo(m[j * 4 + i], 12); + } + }); + + it('reuses its ancestor marker across rows without leaking marks', () => { + // The stamped marker is the one piece of state shared between rows. If a stale generation + // leaked, an LCA would resolve to a node on the PREVIOUS row's path and distances would come + // out too small — so compute a matrix (shared marker) and compare against fresh single rows. + const t = parseNewick('(((A:0.1,B:0.2):0.05,C:0.3):0.01,(D:0.4,E:0.05):0.2);'); + const names = ['A', 'B', 'C', 'D', 'E']; + const nodes = names.map((n) => leafOf(t, n)); + const shared = patristicMatrix(t, nodes); + const rd = rootDistances(t); + for (let i = 0; i < nodes.length; i++) { + const fresh = patristicRow(t, rd, nodes[i], nodes); // no marker -> its own + for (let j = 0; j < nodes.length; j++) { + expect(shared[i * nodes.length + j]).toBeCloseTo(fresh[j], 12); + } + } + }); + + it('propagates a negative branch length into the distance instead of hiding it', () => { + const t = parseNewick('((A:-0.5,B:0.2):0.05,C:0.3);'); + const nodes = ['A', 'B', 'C'].map((n) => leafOf(t, n)); + const m = patristicMatrix(t, nodes); + expect(m[0 * 3 + 1]).toBeCloseTo(-0.3, 12); // A-B: -0.5 + 0.2 + expect(m[0 * 3 + 2]).toBeCloseTo(-0.15, 12); // A-C: -0.5 + 0.05 + 0.3 + }); + + it('gives every pair distance 0 on a topology-only tree', () => { + const t = parseNewick('((A,B),C);'); + const nodes = ['A', 'B', 'C'].map((n) => leafOf(t, n)); + expect(Array.from(patristicMatrix(t, nodes)).every((v) => v === 0)).toBe(true); + }); + + it('handles a deep ladder tree without recursing', () => { + // 3,000 nested clades. The Python reference recurses here and dies at its frame limit; this + // port must not, because DM3 accepts uploads far larger than 1,000 taxa. + const N = 3000; + let s = 'L0:0.001'; + for (let i = 1; i < N; i++) s = `(${s},L${i}:0.001)`; + const t = parseNewick(s + ';'); + const idx = leafIndex(t).index; + expect(idx.size).toBe(N); + const d = rootDistances(t); + // L0 is the deepest tip: N-1 internal branches (all 0, no lengths given) plus its own 0.001. + expect(Number.isFinite(d[idx.get('L0')])).toBe(true); + expect(d[idx.get(`L${N - 1}`)]).toBeCloseTo(0.001, 12); + }); +}); + +describe('maxPdSelect', () => { + it('returns everything, in order, when under the cap', () => { + const t = parseNewick(SIMPLE); + const nodes = ['A', 'B', 'C'].map((n) => leafOf(t, n)); + expect(maxPdSelect(t, nodes, 512).selected).toEqual([0, 1, 2]); + }); + + it('seeds at index 0 and then takes the farthest point', () => { + // A-B are close, C and D are far. Seeded at A (index 0), the next pick is whichever is + // farthest from A, then the one farthest from {A, that}. + const t = parseNewick('((A:0.01,B:0.01):0.05,(C:1.0,D:2.0):0.5);'); + const nodes = ['A', 'B', 'C', 'D'].map((n) => leafOf(t, n)); + const { selected } = maxPdSelect(t, nodes, 3); + expect(selected[0]).toBe(0); // always the seed + expect(selected[1]).toBe(3); // D, farthest from A + expect(selected[2]).toBe(2); // C, farthest from {A, D} + }); + + it('is sensitive to input ORDER, because the seed is index 0 — not a defect this layer fixes', () => { + // The reference always seeds at the first taxon in alignment order, so reordering the + // sequences in an upload changes which taxa the model sees. Pinned so nobody "fixes" it into + // divergence from the model's training-time behaviour. + const t = parseNewick('((A:0.01,B:0.01):0.05,(C:1.0,D:2.0):0.5);'); + const asIs = ['A', 'B', 'C', 'D'].map((n) => leafOf(t, n)); + const rotated = ['C', 'D', 'A', 'B'].map((n) => leafOf(t, n)); + const a = maxPdSelect(t, asIs, 2).selected.map((i) => t.name[asIs[i]]); + const b = maxPdSelect(t, rotated, 2).selected.map((i) => t.name[rotated[i]]); + expect(a).toEqual(['A', 'D']); + expect(b).toEqual(['C', 'D']); + expect(a).not.toEqual(b); + }); + + it('reports duplicates on an all-zero distance matrix instead of hiding them', () => { + // Every selected index has minDist 0, so once all candidates are 0 argmax returns index 0 + // forever and one taxon fills every slot. Reproduced (the model was trained with it) but + // counted, so a caller can refuse. + const t = parseNewick('((A,B),(C,D));'); // no branch lengths -> all distances 0 + const nodes = ['A', 'B', 'C', 'D'].map((n) => leafOf(t, n)); + const { selected, duplicates } = maxPdSelect(t, nodes, 3); + expect(selected).toEqual([0, 0, 0]); + expect(duplicates).toBe(2); + }); + + it('selects exactly maxSpecies when over the cap', () => { + const t = parseNewick('((A:0.1,B:0.2):0.05,(C:0.3,D:0.15):0.02);'); + const nodes = ['A', 'B', 'C', 'D'].map((n) => leafOf(t, n)); + expect(maxPdSelect(t, nodes, 2).selected).toHaveLength(2); + }); +}); diff --git a/src/test/axomeme-postprocess.test.js b/src/test/axomeme-postprocess.test.js new file mode 100644 index 0000000..a0b68f4 --- /dev/null +++ b/src/test/axomeme-postprocess.test.js @@ -0,0 +1,279 @@ +/** + * Tests for AxoMEME output postprocessing. + * + * The transformations here are the ones that turn model outputs into numbers a researcher reads, and + * every one of them is a place where "looks plausible" and "correct" differ by a lot. A missing + * expm1 does not crash; it just reports every rate too low. An invariant site scored rather than + * zeroed does not crash; it reports selection at a site where nothing varies. + * + * Expected values are derived from the reference's arithmetic, not recorded from this code. + */ +import { describe, it, expect } from 'vitest'; +import { + buildPredictions, + isSiteVariable, + siteVariability, + CALL_DEFAULTS, + NEUTRAL_CALL +} from '../lib/services/axomeme/postprocess.js'; + +/** Raw graph outputs for `n` sites, all heads constant unless overridden. */ +function outputs(n, over = {}) { + const fill = (v) => new Float32Array(n).fill(v); + return { + lrt: over.lrt ?? fill(1), + alpha: over.alpha ?? fill(0), + beta_neg: over.beta_neg ?? fill(0), + beta_pos: over.beta_pos ?? fill(0), + p_neg: over.p_neg ?? fill(0.25) + }; +} + +describe('isSiteVariable', () => { + it('is true when more than one amino acid is present', () => { + expect(isSiteVariable(['ATG', 'TTA'])).toBe(true); // M, L + }); + + it('is false when every codon codes the same amino acid', () => { + expect(isSiteVariable(['TTA', 'TTG', 'CTA'])).toBe(false); // all Leucine + expect(isSiteVariable(['ATG', 'ATG'])).toBe(false); + }); + + it('is TRUE for a serine island — synonymous but selection-relevant', () => { + // The condition that is easy to drop. Serine is the one residue whose codons occupy two + // disjoint blocks (TCN and AGY), so switching between them is synonymous yet needs multiple + // substitutions. Every codon here is Serine, so the amino-acid test alone says "invariant". + expect(isSiteVariable(['TCA', 'AGC'])).toBe(true); + expect(isSiteVariable(['TCT', 'AGT'])).toBe(true); + // ...but only when BOTH families are present. + expect(isSiteVariable(['TCA', 'TCG'])).toBe(false); + expect(isSiteVariable(['AGC', 'AGT'])).toBe(false); + }); + + it('is false for an empty site', () => { + expect(isSiteVariable([])).toBe(false); + expect(isSiteVariable(null)).toBe(false); + }); +}); + +describe('siteVariability', () => { + it('ignores gapped and ambiguous codons when judging a site', () => { + // Site 0: ATG / --- / ANT -> only one usable codon, so not variable. + // Site 1: TTA / TTG / TTT -> Leu, Leu, Phe -> variable. + const flags = siteVariability(['ATGTTA', '---TTG', 'ANTTTT'], 2); + expect(flags).toEqual([false, true]); + }); + + it('lets a short sequence contribute nothing past its end', () => { + const flags = siteVariability(['ATGTTA', 'ATG'], 2); + expect(flags[1]).toBe(false); // only one codon reaches site 1 + }); +}); + +describe('buildPredictions', () => { + const sites = (n, variable = true) => ({ + refCodons: new Array(n).fill('ATG'), + variable: new Array(n).fill(variable) + }); + + it('applies expm1 to the rate heads', () => { + // The heads are softplus, so the graph emits log1p(rate). expm1(1) = e - 1 = 1.7182818... + const rows = buildPredictions( + outputs(1, { alpha: new Float32Array([1]), beta_pos: new Float32Array([2]) }), + sites(1) + ); + expect(rows[0].alphaDs).toBeCloseTo(Math.E - 1, 6); + expect(rows[0].betaPosDn).toBeCloseTo(Math.exp(2) - 1, 6); + }); + + it('reports p_pos as 1 - p_neg', () => { + const rows = buildPredictions(outputs(1, { p_neg: new Float32Array([0.25]) }), sites(1)); + expect(rows[0].pPos).toBeCloseTo(0.75, 6); + }); + + it('treats lrt as the LRT and DERIVES the log, not the other way round', () => { + // Worth pinning: the export is eval mode, so `lrt` is already ordinal-decoded and is the LRT + // itself. The reference computes predicted_log_lrt = log1p(predicted_lrt). + const rows = buildPredictions(outputs(1, { lrt: new Float32Array([4]) }), sites(1)); + expect(rows[0].lrt).toBeCloseTo(4, 6); + expect(rows[0].logLrt).toBeCloseTo(Math.log1p(4), 6); + }); + + it('clamps a negative lrt or rate to zero', () => { + const rows = buildPredictions( + outputs(1, { lrt: new Float32Array([-2]), alpha: new Float32Array([-3]) }), + sites(1) + ); + expect(rows[0].lrt).toBe(0); + expect(rows[0].alphaDs).toBe(0); + }); + + it('ZEROES an invariant site instead of scoring it', () => { + // The reference zeroes before consulting the model. A large model output at an invariant site + // must not reach the user — those zeros mean "not applicable", not "no selection". + const rows = buildPredictions(outputs(2, { lrt: new Float32Array([99, 99]) }), { + refCodons: ['ATG', 'ATG'], + variable: [false, true] + }); + expect(rows[0].lrt).toBe(0); + expect(rows[0].call).toBe(NEUTRAL_CALL); + expect(rows[0].isVariable).toBe(false); + expect(rows[1].lrt).toBe(99); + }); + + it('excludes invariant sites from the local statistics', () => { + // If the zeroed sites were included they would drag the mean down and inflate every z-score. + const lrt = new Float32Array([0, 10, 20, 30]); + const withInvariant = buildPredictions(outputs(4, { lrt }), { + refCodons: new Array(4).fill('ATG'), + variable: [false, true, true, true] + }); + const onlyVariable = buildPredictions(outputs(3, { lrt: new Float32Array([10, 20, 30]) }), { + refCodons: new Array(3).fill('ATG'), + variable: [true, true, true] + }); + expect(withInvariant[1].zScore).toBeCloseTo(onlyVariable[0].zScore, 10); + expect(withInvariant[3].zScore).toBeCloseTo(onlyVariable[2].zScore, 10); + }); + + it('computes z-scores with the POPULATION standard deviation', () => { + // np.std, not pandas' sample std — the difference is a factor of sqrt(n/(n-1)), which for + // three sites is 22%. + const rows = buildPredictions(outputs(3, { lrt: new Float32Array([1, 2, 3]) }), sites(3)); + // mean 2, population std = sqrt(2/3) = 0.8165 + expect(rows[0].zScore).toBeCloseTo(-1 / Math.sqrt(2 / 3), 6); + expect(rows[1].zScore).toBeCloseTo(0, 10); + expect(rows[2].zScore).toBeCloseTo(1 / Math.sqrt(2 / 3), 6); + }); + + it('gives a zero z-score when every site is identical, rather than dividing by zero', () => { + const rows = buildPredictions(outputs(3, { lrt: new Float32Array([5, 5, 5]) }), sites(3)); + for (const r of rows) expect(r.zScore).toBe(0); + }); + + it('averages percentile ranks across ties, matching pandas', () => { + // [10, 20, 20, 40]: the tied pair share rank (2+3)/2 = 2.5 -> 62.5%. + const rows = buildPredictions( + outputs(4, { lrt: new Float32Array([10, 20, 20, 40]) }), + sites(4) + ); + expect(rows[0].percentile).toBeCloseTo(25, 6); + expect(rows[1].percentile).toBeCloseTo(62.5, 6); + expect(rows[2].percentile).toBeCloseTo(62.5, 6); + expect(rows[3].percentile).toBeCloseTo(100, 6); + }); + + it("DEFAULTS TO percentile, not the reference driver's pvalue", () => { + // Measured, not preferential. Across 12 real DataMonkey submissions and 662 variable sites the + // highest predicted LRT anywhere was 3.902; one site cleared the 3.12 gate and none cleared + // 4.45 — including an alignment where MEME itself reports 17 sites at p <= 0.05. Under pvalue + // this feature reports nothing on real data. + expect(CALL_DEFAULTS.mode).toBe('percentile'); + // The gates themselves are unchanged, so switching modes still reproduces the reference. + expect(CALL_DEFAULTS.tier1LrtGate).toBe(4.45); + expect(CALL_DEFAULTS.tier2LrtGate).toBe(3.12); + expect(CALL_DEFAULTS.tier1Percentile).toBe(98.0); + expect(CALL_DEFAULTS.tier2Percentile).toBe(95.0); + }); + + it('returns top-ranked sites by default even when no score approaches an LRT gate', () => { + // The whole reason for the default. These scores are realistic for the model — nothing near + // 3.12 — and pvalue mode calls nothing on them while percentile surfaces the top of the range. + const lrt = new Float32Array(100); + for (let i = 0; i < 100; i++) lrt[i] = (i / 99) * 2.5; + const ranked = buildPredictions(outputs(100, { lrt }), sites(100)); + expect(ranked.filter((r) => r.call !== NEUTRAL_CALL).length).toBeGreaterThan(0); + const gated = buildPredictions(outputs(100, { lrt }), sites(100), { mode: 'pvalue' }); + expect(gated.filter((r) => r.call !== NEUTRAL_CALL)).toEqual([]); + }); + + it('labels tiers by what they MEAN, not by a confidence word', () => { + // "High"/"Medium" imply a calibrated confidence the model does not have. Each label states the + // rule that produced it, so a reader can see it is relative to this alignment. + const lrt = new Float32Array(100); + for (let i = 0; i < 100; i++) lrt[i] = i; + const pct = buildPredictions(outputs(100, { lrt }), sites(100)); + expect(pct.map((r) => r.call)).toContain('Top 2%'); + expect(pct.map((r) => r.call)).toContain('Top 5%'); + + // A uniform spread cannot reach z = 2.5 — its maximum z is about 1.73 regardless of n — so this + // needs a genuine outlier. Same bound as the short-alignment case pinned further down. + const spike = new Float32Array(100).fill(1); + spike[42] = 100; + const z = buildPredictions(outputs(100, { lrt: spike }), sites(100), { mode: 'zscore' }); + expect(z[42].call).toBe('Z ≥ 2.5'); + + const p = buildPredictions( + outputs(4, { lrt: new Float32Array([5.0, 3.5, 3.13, 1.0]) }), + sites(4), + { mode: 'pvalue' } + ); + expect(p[0].call).toBe('LRT ≥ 4.45'); + expect(p[1].call).toBe('LRT ≥ 3.12'); + expect(p[3].call).toBe(NEUTRAL_CALL); + }); + + it('still reproduces the reference exactly when pvalue mode is chosen', () => { + const rows = buildPredictions( + outputs(4, { lrt: new Float32Array([5.0, 3.5, 3.13, 1.0]) }), + sites(4), + { mode: 'pvalue' } + ); + expect(rows[0].call).toMatch(/4\.45/); + expect(rows[1].call).toMatch(/3\.12/); + expect(rows[2].call).toMatch(/3\.12/); + expect(rows[3].call).toBe(NEUTRAL_CALL); + }); + + it('does NOT call a site whose float32 LRT lands just under the gate', () => { + // A boundary worth pinning rather than smoothing. The gates are float64 literals but the model + // emits float32, and 3.12 is not representable: Float32Array([3.12])[0] is 3.119999885559082, + // which fails `>= 3.12`. The reference behaves identically — torch's .item() widens the same + // float32 to float64 — so this is faithful, not a rounding bug to fix. It does mean a site + // sitting exactly on a gate falls to the lower tier. + const exact = new Float32Array([3.12]); + expect(exact[0]).toBeLessThan(3.12); + const rows = buildPredictions(outputs(1, { lrt: exact }), sites(1), { mode: 'pvalue' }); + expect(rows[0].call).toBe(NEUTRAL_CALL); + }); + + it('supports the zscore and percentile call modes', () => { + // NOTE the site count. With a POPULATION standard deviation, |z| is bounded by sqrt(n-1), so + // zscore mode cannot call anything at all on fewer than 5 variable sites — at n=4 the maximum + // possible z is 1.732, below even the Tier 2 threshold of 2.0. Ten sites clears it. + const lrt = new Float32Array([1, 1, 1, 1, 1, 1, 1, 1, 1, 100]); + const z = buildPredictions(outputs(10, { lrt }), sites(10), { mode: 'zscore' }); + expect(z[9].call).toMatch(/^Z ≥ /); // the outlier + expect(z[0].call).toBe(NEUTRAL_CALL); + + const p = buildPredictions(outputs(10, { lrt }), sites(10), { + mode: 'percentile', + tier1Percentile: 90, + tier2Percentile: 70 + }); + expect(p[9].call).toBe('Top 10%'); + expect(p[0].call).toBe(NEUTRAL_CALL); + }); + + it('zscore mode is structurally unable to call short alignments', () => { + // The consequence of the sqrt(n-1) bound, stated as its own fact so it is not rediscovered as + // "the model found nothing". Four variable sites, one of them enormous, and nothing calls. + const rows = buildPredictions( + outputs(4, { lrt: new Float32Array([1, 2, 3, 1000]) }), + sites(4), + { mode: 'zscore' } + ); + expect(Math.max(...rows.map((r) => r.zScore))).toBeLessThan(Math.sqrt(3) + 1e-9); + expect(rows.every((r) => r.call === NEUTRAL_CALL)).toBe(true); + }); + + it('does not invent an amino acid for a gapped reference codon', () => { + const rows = buildPredictions(outputs(1), { refCodons: ['A-G'], variable: [true] }); + expect(rows[0].refAa).toBe('?'); + }); + + it('numbers sites from 1', () => { + const rows = buildPredictions(outputs(3), sites(3)); + expect(rows.map((r) => r.site)).toEqual([1, 2, 3]); + }); +}); diff --git a/src/test/axomeme-review-regressions.test.js b/src/test/axomeme-review-regressions.test.js new file mode 100644 index 0000000..2781d12 --- /dev/null +++ b/src/test/axomeme-review-regressions.test.js @@ -0,0 +1,35 @@ +import { describe, it, expect } from 'vitest'; +import { prepareAlignment, batchSizeFor } from '../lib/services/axomeme/assemble.js'; + +describe('review regressions', () => { + it('exposes indices, so a duplicate FASTA header cannot resolve to the wrong record', () => { + // names.indexOf() returns the FIRST match; orderSpecies keeps the LAST. Without indices the + // variability flags would be computed from a different sequence than the model was fed. + const names = ['dup', 'other', 'dup']; + const sequences = ['ATGAAA', 'ATGTTT', 'ATGCCC']; + const p = prepareAlignment({ + names, + sequences, + treeText: '((dup:0.1,other:0.2):0.05,dup:0.3);', + maxSpecies: 8 + }); + expect(p.selectedIndices).toBeDefined(); + expect(p.selectedIndices.every((i) => Number.isInteger(i))).toBe(true); + // The duplicate must resolve to an index whose name really is that name... + for (let k = 0; k < p.selectedIndices.length; k++) { + expect(names[p.selectedIndices[k]]).toBe(p.selectedNames[k]); + } + // ...and at least one selected index must NOT be the naive indexOf answer, or this alignment + // would not exercise the bug. + const naive = p.selectedNames.map((n) => names.indexOf(n)); + expect(naive).not.toEqual(p.selectedIndices); + expect(Number.isInteger(p.referenceIndex)).toBe(true); + }); + + it('batchSizeFor can exceed the argument-spread limit, so results must not be spread', () => { + // The number that made `push(...batch)` throw "Maximum call stack size exceeded" at ~90% + // progress. Pinned so nobody reintroduces the spread thinking the batches are small. + expect(batchSizeFor(10)).toBeGreaterThan(125000); + expect(batchSizeFor(4)).toBeGreaterThan(125000); + }); +}); diff --git a/src/test/axomeme-session.test.js b/src/test/axomeme-session.test.js new file mode 100644 index 0000000..6867c66 --- /dev/null +++ b/src/test/axomeme-session.test.js @@ -0,0 +1,197 @@ +/** + * Tests for AxoMEME session loading. + * + * The load path is where 17 MB of runtime and model either does or does not reach a user, and where a + * swapped artifact either is or is not caught. Both failure modes are silent — the wrong model still + * returns five plausible numbers per site — so every guard here is exercised in both directions. + * + * The ONNX runtime is faked. These tests are about the loading policy, not about inference; the real + * graph is exercised by the end-to-end check, which needs the 3.78 MB artifact and a browser. + */ +import { describe, it, expect, beforeEach, vi } from 'vitest'; +import { + loadSession, + resetSession, + isSessionLoaded, + runSites, + MODEL_URL +} from '../lib/services/axomeme/session.js'; +import { VERIFIED_MODEL_SHA256 } from '../lib/services/axomeme/modelContract.js'; + +/** Bytes whose sha256 is the pinned value — impossible to forge, so tests bypass the check instead. */ +const SOME_BYTES = new Uint8Array([1, 2, 3, 4]).buffer; + +function fakeOrt( + inputNames = ['msa_codons', 'msa_aas', 'dist_matrix', 'mds_coords', 'padding_mask'] +) { + const created = []; + return { + created, + Tensor: class { + constructor(type, data, dims) { + this.type = type; + this.data = data; + this.dims = dims; + } + }, + InferenceSession: { + create: vi.fn(async (bytes) => { + created.push(bytes.byteLength ?? bytes.length); + return { + inputNames, + outputNames: ['lrt', 'alpha', 'beta_neg', 'beta_pos', 'p_neg'], + run: vi.fn(async (feeds) => { + const n = feeds.msa_codons.dims[0]; + const f = () => ({ data: new Float32Array(n).fill(1) }); + return { lrt: f(), alpha: f(), beta_neg: f(), beta_pos: f(), p_neg: f() }; + }) + }; + }) + } + }; +} + +const okFetch = (buffer = SOME_BYTES) => + vi.fn(async () => ({ ok: true, status: 200, statusText: 'OK', arrayBuffer: async () => buffer })); + +describe('loadSession', () => { + beforeEach(() => resetSession()); + + it('does not load anything just by importing the module', () => { + // The whole cost discipline in one assertion: fourteen other methods import nothing by + // touching this file. + expect(isSessionLoaded()).toBe(false); + }); + + it('fetches the model from static/, not from the bundle', () => { + // A bundled model would be a build-time import and would land in the JS graph. + expect(MODEL_URL).toMatch(/^\/models\//); + expect(MODEL_URL).toMatch(/\.onnx$/); + }); + + it('loads and reports the byte count', async () => { + const ort = fakeOrt(); + const r = await loadSession({ ort, fetchImpl: okFetch(), verifyHash: false }); + expect(r.bytes).toBe(4); + expect(ort.InferenceSession.create).toHaveBeenCalledTimes(1); + }); + + it('REFUSES a model whose hash is not the verified one', async () => { + // The contract's conclusions — eval-mode export, lrt already decoded, rate heads in log1p + // space — were read off one specific graph. A different one may be fine; nothing here knows + // that, and guessing produces wrong numbers rather than an error. + const err = await loadSession({ ort: fakeOrt(), fetchImpl: okFetch() }).catch((e) => e); + expect(err).toBeInstanceOf(Error); + expect(err.message).toMatch(/hash mismatch/); + expect(err.message).toContain(VERIFIED_MODEL_SHA256); + // And it must say what to do about it, not just that it failed. + expect(err.message).toMatch(/VERIFIED_MODEL_SHA256/); + }); + + it('refuses a graph that is missing a contract input', async () => { + const ort = fakeOrt(['msa_codons', 'msa_aas', 'dist_matrix']); // no mds_coords / padding_mask + const err = await loadSession({ ort, fetchImpl: okFetch(), verifyHash: false }).catch((e) => e); + expect(err.message).toMatch(/missing expected inputs/); + expect(err.message).toMatch(/mds_coords/); + expect(err.message).toMatch(/padding_mask/); + }); + + it('reports a failed fetch with its status', async () => { + const fetchImpl = vi.fn(async () => ({ ok: false, status: 404, statusText: 'Not Found' })); + const err = await loadSession({ ort: fakeOrt(), fetchImpl, verifyHash: false }).catch((e) => e); + expect(err.message).toMatch(/404/); + }); + + it('memoises the PRODUCTION path, so a second alignment does not re-download 17 MB', async () => { + // This test used to pass `ort`/`fetchImpl`, which loadSession treats as "do not memoise" — so it + // asserted the BYPASS path and the memo itself had no coverage at all. A regression that dropped + // the memo would have left the suite green while every alignment re-downloaded 3.78 MB of model. + // + // The production path takes no options, so it is driven here by stubbing the globals it reaches + // for. onnxruntime-web cannot be imported under jsdom, so the assertion is on the FETCH count: + // one download no matter how many callers ask. + const buffer = SOME_BYTES; + let fetches = 0; + const realFetch = globalThis.fetch; + globalThis.fetch = async () => { + fetches++; + return { ok: true, status: 200, statusText: 'OK', arrayBuffer: async () => buffer }; + }; + try { + const a = loadSession(); + const b = loadSession(); + expect(a, 'a second call returned a different promise — the memo is not shared').toBe(b); + expect(isSessionLoaded(), 'isSessionLoaded stayed false on the production path').toBe(true); + await Promise.allSettled([a, b]); + // It fails at the hash check or the ort import, but only ONE fetch may have happened. + expect(fetches).toBeLessThanOrEqual(1); + } finally { + globalThis.fetch = realFetch; + } + }); + + it('bypasses the memo when options are supplied, so tests stay independent', async () => { + const ort = fakeOrt(); + const fetchImpl = okFetch(); + await loadSession({ ort, fetchImpl, verifyHash: false }); + await loadSession({ ort, fetchImpl, verifyHash: false }); + expect(fetchImpl).toHaveBeenCalledTimes(2); + // And an options call must never populate the production memo — including verifyHash:false, + // which would otherwise cache a session that was never checked against the pinned artifact. + expect(isSessionLoaded()).toBe(false); + }); + + it('does not memoise a FAILURE, so a transient error is recoverable', async () => { + // Caching a rejected promise would disable the feature for the life of the page. + await loadSession({ ort: fakeOrt(), fetchImpl: okFetch() }).catch(() => {}); + expect(isSessionLoaded()).toBe(false); + }); +}); + +describe('runSites', () => { + it('feeds every tensor with the dtype the graph expects', async () => { + const ort = fakeOrt(); + const { session } = await loadSession({ ort, fetchImpl: okFetch(), verifyHash: false }); + const B = 2; + const N = 3; + const bundle = { + msa_codons: { data: new BigInt64Array(B * N), dims: [B, N, 1] }, + msa_aas: { data: new BigInt64Array(B * N), dims: [B, N, 1] }, + dist_matrix: { data: new Float32Array(B * N * N), dims: [B, N, N] }, + mds_coords: { data: new Float32Array(B * N * 4), dims: [B, N, 4] }, + padding_mask: { data: new Uint8Array(B * N), dims: [B, N] } + }; + const out = await runSites(session, bundle, ort); + const feeds = session.run.mock.calls[0][0]; + // int64 for the token streams: passing float32 of the same values is a type error at + // session.run, not a silent coercion. + expect(feeds.msa_codons.type).toBe('int64'); + expect(feeds.msa_aas.type).toBe('int64'); + expect(feeds.dist_matrix.type).toBe('float32'); + expect(feeds.mds_coords.type).toBe('float32'); + expect(feeds.padding_mask.type).toBe('bool'); + expect(Object.keys(out).sort()).toEqual(['alpha', 'beta_neg', 'beta_pos', 'lrt', 'p_neg']); + expect(out.lrt).toHaveLength(B); + }); + + it('runs ALL sites in one call, not one call per site', async () => { + // The three phylo tensors are per-alignment, so batching is nearly free. The reference driver + // loops one forward pass per codon; for a 441-site alignment that is 441x the overhead. + const ort = fakeOrt(); + const { session } = await loadSession({ ort, fetchImpl: okFetch(), verifyHash: false }); + const B = 441; + const N = 2; + await runSites( + session, + { + msa_codons: { data: new BigInt64Array(B * N), dims: [B, N, 1] }, + msa_aas: { data: new BigInt64Array(B * N), dims: [B, N, 1] }, + dist_matrix: { data: new Float32Array(B * N * N), dims: [B, N, N] }, + mds_coords: { data: new Float32Array(B * N * 4), dims: [B, N, 4] }, + padding_mask: { data: new Uint8Array(B * N), dims: [B, N] } + }, + ort + ); + expect(session.run).toHaveBeenCalledTimes(1); + }); +}); diff --git a/src/test/axomeme-tokenizer.test.js b/src/test/axomeme-tokenizer.test.js new file mode 100644 index 0000000..77e5907 --- /dev/null +++ b/src/test/axomeme-tokenizer.test.js @@ -0,0 +1,184 @@ +/** + * Tests for AxoMEME codon/amino-acid tokenisation. + * + * The block that matters most here is the last one. This port deliberately implements the TRAINING + * tokenizer and deliberately does NOT implement the one in the handoff's inference driver, because + * those two disagree on 63 of 64 codons. That is the kind of decision a future reader "corrects" in + * good faith — the driver is, after all, the script the ML team ships — so the divergence is pinned + * from both sides: these tests fail if we drift away from training, AND they fail if someone aligns + * us to the driver. + * + * Everything else is the ordinary business of a tokenizer, with one thing worth stating: the + * expected translations are the standard genetic code, checkable against any codon table, not values + * recorded from this implementation. + */ +import { describe, it, expect } from 'vitest'; +import { + CODON_LIST, + CODON_TO_IDX, + GENETIC_CODE, + AA_TO_IDX, + codonToken, + aaToken, + tokenizeSequence +} from '../lib/services/axomeme/tokenizer.js'; +import { + CODON_GAP, + CODON_UNKNOWN, + AA_GAP, + AA_UNKNOWN +} from '../lib/services/axomeme/modelContract.js'; + +describe('the codon vocabulary', () => { + it('is the 64 codons in TCAG order', () => { + expect(CODON_LIST).toHaveLength(64); + expect(new Set(CODON_LIST).size).toBe(64); + expect(CODON_LIST[0]).toBe('TTT'); + expect(CODON_LIST[63]).toBe('GGG'); + // Third position varies fastest, matching [a][b][c] with c innermost. + expect(CODON_LIST.slice(0, 4)).toEqual(['TTT', 'TTC', 'TTA', 'TTG']); + }); + + it('includes stop codons, which are real tokens and not sentinels', () => { + for (const stop of ['TAA', 'TAG', 'TGA']) { + expect(CODON_TO_IDX.has(stop)).toBe(true); + expect(codonToken(stop)).toBeLessThan(64); + } + }); +}); + +describe('the genetic code', () => { + it('translates the standard code', () => { + // Spot values anyone can check against a codon table. + const expected = { + ATG: 'M', // Met / start + TGG: 'W', // the other single-codon residue + TTT: 'F', + TTA: 'L', + AAA: 'K', + GGG: 'G', + TAA: '*', + TAG: '*', + TGA: '*' + }; + for (const [codon, aa] of Object.entries(expected)) { + expect(GENETIC_CODE.get(codon), codon).toBe(aa); + } + }); + + it('covers every codon and has exactly three stops', () => { + expect(GENETIC_CODE.size).toBe(64); + for (const c of CODON_LIST) expect(GENETIC_CODE.get(c), c).toBeTruthy(); + expect( + [...GENETIC_CODE.entries()] + .filter(([, a]) => a === '*') + .map(([c]) => c) + .sort() + ).toEqual(['TAA', 'TAG', 'TGA']); + }); + + it('uses "*" for a stop, so it lands on a real AA token rather than unknown', () => { + // The handoff's driver writes stops as '_' while its AA_LIST contains '*', so stops fall + // through to 22 (unknown) there. Here they translate to 20. + expect(AA_TO_IDX.get('*')).toBe(20); + expect(aaToken('TAA')).toBe(20); + expect(aaToken('TAA')).not.toBe(AA_UNKNOWN); + }); +}); + +describe('codonToken', () => { + it('is case insensitive', () => { + expect(codonToken('atg')).toBe(codonToken('ATG')); + }); + + it('treats a gap ANYWHERE as a gap, not as unknown', () => { + // `'-' in codon` is a substring test in the reference, so a partial gap counts. + expect(codonToken('---')).toBe(CODON_GAP); + expect(codonToken('A-T')).toBe(CODON_GAP); + expect(codonToken('AT-')).toBe(CODON_GAP); + // And a gap outranks the length and N tests: this is length 2 AND gapped. + expect(codonToken('A-')).toBe(CODON_GAP); + }); + + it('maps ambiguity and wrong lengths to unknown', () => { + expect(codonToken('ANT')).toBe(CODON_UNKNOWN); + expect(codonToken('NNN')).toBe(CODON_UNKNOWN); + expect(codonToken('AT')).toBe(CODON_UNKNOWN); + expect(codonToken('ATGC')).toBe(CODON_UNKNOWN); + expect(codonToken('')).toBe(CODON_UNKNOWN); + expect(codonToken('XYZ')).toBe(CODON_UNKNOWN); + }); +}); + +describe('aaToken', () => { + it('translates, and separates gap from unknown', () => { + expect(aaToken('ATG')).toBe(AA_TO_IDX.get('M')); + expect(aaToken('---')).toBe(AA_GAP); + expect(aaToken('A-T')).toBe(AA_GAP); + expect(aaToken('ANT')).toBe(AA_UNKNOWN); + expect(aaToken('AT')).toBe(AA_UNKNOWN); + expect(AA_GAP).not.toBe(AA_UNKNOWN); + }); + + it('rejects "?" where codonToken does not — the reference asymmetry, same answer either way', () => { + expect(aaToken('?A?')).toBe(AA_UNKNOWN); + // codonToken has no '?' test, but a '?' codon is not in the map, so it lands on unknown too. + expect(codonToken('?A?')).toBe(CODON_UNKNOWN); + }); +}); + +describe('tokenizeSequence', () => { + it('tokenises in frame and drops a trailing partial codon', () => { + // 7 nt -> 2 whole codons, matching len // 3. + const { codons, aas } = tokenizeSequence('ATGTTTAA'); + expect(codons).toHaveLength(2); + expect(Array.from(codons)).toEqual([codonToken('ATG'), codonToken('TTT')]); + expect(Array.from(aas)).toEqual([AA_TO_IDX.get('M'), AA_TO_IDX.get('F')]); + }); + + it('handles an empty and a sub-codon sequence', () => { + expect(tokenizeSequence('').codons).toHaveLength(0); + expect(tokenizeSequence('AT').codons).toHaveLength(0); + }); + + it('carries gaps through as gap tokens', () => { + const { codons, aas } = tokenizeSequence('ATG---'); + expect(Array.from(codons)).toEqual([codonToken('ATG'), CODON_GAP]); + expect(Array.from(aas)).toEqual([AA_TO_IDX.get('M'), AA_GAP]); + }); +}); + +describe('WE USE THE TRAINING TOKENIZER, NOT THE INFERENCE DRIVER"S', () => { + /** + * Measured over all 64 codons: predict_regression_nexus.py builds a 60-codon ALPHABETICAL + * vocabulary (lines 49-55), then redefines get_codon_token (line 74) WITHOUT redefining + * CODON_TO_IDX, and imports only three non-tokenizer names from the training module. So the + * driver disagrees with training on 63 of 64 codons, and drops TTA plus all three stops entirely. + * + * A model trained on TCAG-64 must be served TCAG-64. These are the exact values that distinguish + * the two, so this block fails whichever way someone drifts. + */ + const TRAINING = { TTT: 0, TTA: 2, TAA: 10, ATG: 35, AAA: 42, GGG: 63 }; + const DRIVER = { TTT: 59, TTA: 65, TAA: 65, ATG: 14, AAA: 0, GGG: 42 }; + + it('matches the training vocabulary exactly', () => { + for (const [codon, token] of Object.entries(TRAINING)) { + expect(codonToken(codon), codon).toBe(token); + } + }); + + it('does NOT match the driver vocabulary', () => { + for (const [codon, driverToken] of Object.entries(DRIVER)) { + expect(codonToken(codon), `${codon} matched the driver's token`).not.toBe(driverToken); + } + }); + + it('keeps TTA and the stops as real codons, which the driver discards', () => { + // The driver's CODON_LIST omits TTA (Leucine — almost certainly a slip when the stops were + // removed) along with TAA/TAG/TGA, so all four collapse to "unknown" there and real leucine + // data is thrown away. + for (const c of ['TTA', 'TAA', 'TAG', 'TGA']) { + expect(codonToken(c), c).toBeLessThan(64); + } + }); +}); diff --git a/src/test/meme-hit-likelihood.test.js b/src/test/meme-hit-likelihood.test.js index 261fe08..36c6ab3 100644 --- a/src/test/meme-hit-likelihood.test.js +++ b/src/test/meme-hit-likelihood.test.js @@ -406,7 +406,11 @@ describe('the walker refuses a malformed model', () => { const REFUSALS = [ // --- multi-class: the walker sums ONE tree group; K groups produce a plausible number with // --- no meaning at all. - ['multi-class (declared)', patch(modelDoc, [...LMP, 'num_class'], '3'), /num_class=3|multi-class/i], + [ + 'multi-class (declared)', + patch(modelDoc, [...LMP, 'num_class'], '3'), + /num_class=3|multi-class/i + ], [ 'multi-class (structural, via tree_info)', patch( @@ -429,10 +433,18 @@ describe('the walker refuses a malformed model', () => { ), /categorical/i ], - ['categorical split (categories_nodes)', patch(modelDoc, [...TREE0, 'categories_nodes'], [0]), /categorical/i], + [ + 'categorical split (categories_nodes)', + patch(modelDoc, [...TREE0, 'categories_nodes'], [0]), + /categorical/i + ], [ 'categorical encoding table (cats.enc)', - patch(modelDoc, [...GBM, 'cats'], { enc: [{ values: [1, 2, 3] }], feature_segments: [], sorted_idx: [] }), + patch(modelDoc, [...GBM, 'cats'], { + enc: [{ values: [1, 2, 3] }], + feature_segments: [], + sorted_idx: [] + }), /categorical/i ], @@ -480,7 +492,11 @@ describe('the walker refuses a malformed model', () => { patch(modelDoc, [...GBM, 'gbtree_model_param', 'num_parallel_tree'], '4'), /num_parallel_tree|forest/i ], - ['vector leaves', patch(modelDoc, [...TREE0, 'tree_param', 'size_leaf_vector'], '2'), /size_leaf_vector|vector leaves/i], + [ + 'vector leaves', + patch(modelDoc, [...TREE0, 'tree_param', 'size_leaf_vector'], '2'), + /size_leaf_vector|vector leaves/i + ], // --- topology: a child index past the end reads `undefined` and compares false forever; a // --- child index of 0 is a cycle back to the root, which would hang the browser. @@ -489,7 +505,9 @@ describe('the walker refuses a malformed model', () => { patch( modelDoc, [...TREE0, 'right_children'], - modelDoc.learner.gradient_booster.model.trees[0].right_children.map((r, i) => (i === 0 ? 999999 : r)) + modelDoc.learner.gradient_booster.model.trees[0].right_children.map((r, i) => + i === 0 ? 999999 : r + ) ), /out of range/i ], @@ -498,7 +516,9 @@ describe('the walker refuses a malformed model', () => { patch( modelDoc, [...TREE0, 'left_children'], - modelDoc.learner.gradient_booster.model.trees[0].left_children.map((l, i) => (i === 0 ? 0 : l)) + modelDoc.learner.gradient_booster.model.trees[0].left_children.map((l, i) => + i === 0 ? 0 : l + ) ), /out of range/i ], @@ -507,7 +527,9 @@ describe('the walker refuses a malformed model', () => { patch( modelDoc, [...TREE0, 'right_children'], - modelDoc.learner.gradient_booster.model.trees[0].right_children.map((r, i) => (i === 0 ? -1 : r)) + modelDoc.learner.gradient_booster.model.trees[0].right_children.map((r, i) => + i === 0 ? -1 : r + ) ), /one child set to -1/i ], @@ -528,7 +550,11 @@ describe('the walker refuses a malformed model', () => { /written by XGBoost 1\.7\.6/ ], ['no version array at all', patch(modelDoc, ['version'], undefined), /version/i], - ['a base_score the walker will not guess at', patch(modelDoc, [...LMP, 'base_score'], '[0.3,0.3,0.4]'), /components/i] + [ + 'a base_score the walker will not guess at', + patch(modelDoc, [...LMP, 'base_score'], '[0.3,0.3,0.4]'), + /components/i + ] ]; for (const [name, doc, pattern] of REFUSALS) { @@ -1187,7 +1213,10 @@ describe('the copy states a rate, and states the same rate everywhere', () => { expect(stated, `the uncertain band no longer quotes an interval: ${sentence}`).not.toBeNull(); const p = hit / n; const half = 1.96 * Math.sqrt((p * (1 - p)) / n); - for (const [i, want] of [Math.round(100 * (p - half)), Math.round(100 * (p + half))].entries()) { + for (const [i, want] of [ + Math.round(100 * (p - half)), + Math.round(100 * (p + half)) + ].entries()) { expect( Number(stated[i + 1]), `the copy quotes [${stated[1]}%, ${stated[2]}%]; ${hit}/${n} gives ` + @@ -1480,7 +1509,11 @@ describe('scope: the estimate only ever talks about MEME', () => { ]; const levelsSeen = new Set(); for (const [name, args] of states) { - const res = await estimateHitLikelihood({ ...args, model, opts: { resampleAvailable: true } }); + const res = await estimateHitLikelihood({ + ...args, + model, + opts: { resampleAvailable: true } + }); if (res.level) levelsSeen.add(res.level); const text = visibleText(res); expect(text, `state "${name}" named another method: ${text}`).not.toMatch(OTHER_METHODS); @@ -1533,20 +1566,84 @@ describe('scope: the estimate only ever talks about MEME', () => { expect(MODEL_BASIS).toMatch(/MEME/); }); - it('ships no ML runtime, in package.json or anywhere the estimator can reach', () => { - // A tombstone. This feature's first version downloaded 13.5 MB of ONNX Runtime WASM for - // every method, including the ~14 that render no estimate. The whole conversion pipeline is - // gone — the browser parses XGBoost's own save_model() output — and these two assertions are - // what stop a runtime coming back with a future model. + it('reaches no ML runtime from anywhere in the gate', () => { + // A tombstone, RESCOPED — read this before you touch it. + // + // This assertion was written as "DM3 ships no ML runtime, anywhere", and for as long as the + // gate was the only model in the repository those two statements were the same statement. + // They are not any more. AxoMEME 2.0 is a 3.78 MB transformer that genuinely needs + // onnxruntime-web: its graph is a real neural network, not 500 trees of three features, and + // no 40-line walker is going to execute it. So the repo-wide ban would now fail for a + // legitimate reason, and the fix someone reaches for when a guard fails legitimately is + // deleting the guard. + // + // What was actually being defended was never "no runtime exists". It was: THE GATE COSTS + // ALMOST NOTHING AND IS REACHABLE FROM EVERY METHOD, SO IT MUST NOT DRAG A RUNTIME BEHIND + // IT. The gate renders for one method out of fifteen; the first version of this feature + // downloaded 13.5 MB of ONNX Runtime WASM for all fifteen and then rendered nothing for + // fourteen of them. That failure mode gets MORE available once a runtime is a legitimate + // dependency, not less, because now a stray import resolves instead of erroring. + // + // So the guard is now about REACHABILITY, which is the property that was always load-bearing: + // walk the gate's own import closure and prove no ML runtime is in it. That is strictly + // stronger than the per-file grep it replaces — the old version listed four files by hand and + // would not have noticed a fifth. + const RUNTIME = /^(onnxruntime|@tensorflow\/|@xenova\/|onnx|torch|tflite)/i; + // THREE forms, because the first version of this test had only the first and therefore did + // not fire when `import 'onnxruntime-web';` was injected into a reachable module to prove it + // could. A side-effect import has no `from` clause, and it is the exact shape a runtime + // arrives in — you import it for the WASM registration, not for a binding. + const SPECIFIER_FORMS = [ + /\bfrom\s*['"]([^'"]+)['"]/g, // import … from 'x' / export … from 'x' + /\bimport\s*\(\s*['"]([^'"]+)['"]\s*\)/g, // dynamic import('x') + /^\s*import\s+['"]([^'"]+)['"]/gm // side-effect import 'x' + ]; + + const seen = new Set(); + const offenders = []; + const visit = (absPath) => { + if (seen.has(absPath) || !existsSync(absPath)) return; + seen.add(absPath); + // Strip comments before matching. hitLikelihoodModel.js documents its own usage with a + // literal `import … from './hitLikelihood.js'` inside a doc block, and a commented-out + // runtime import must not be reported as a live edge in either direction. + const src = readFileSync(absPath, 'utf8') + .replace(/\/\*[\s\S]*?\*\//g, '') + .replace(/^\s*\/\/.*$/gm, ''); + for (const re of SPECIFIER_FORMS) { + re.lastIndex = 0; + let m; + while ((m = re.exec(src)) !== null) { + const spec = m[1]; + if (!spec) continue; + if (RUNTIME.test(spec)) { + offenders.push(`${absPath.replace(REPO, '')} imports ${spec}`); + continue; + } + // Follow relative edges only. A bare specifier that is not a runtime is somebody + // else's package and cannot pull one in without appearing in package.json, which + // is checked separately below. + if (spec.startsWith('.')) visit(join(dirname(absPath), spec)); + } + } + }; + // Every entry point a caller can reach the gate through. + for (const f of ['scope.js', 'hitLikelihood.js', 'hitLikelihoodModel.js', 'xgbEnsemble.js']) { + visit(join(PRESCREEN, f)); + } + expect(offenders, 'an ML runtime is reachable from the gate').toEqual([]); + // The walk has to have actually walked. Without this the test passes trivially if the entry + // filenames are ever renamed out from under it. + expect(seen.size).toBeGreaterThanOrEqual(4); + + // package.json is still checked, but as an ALLOWLIST rather than a ban. onnxruntime-web is + // permitted because AxoMEME needs it; anything else in this family is not, and adding one is + // now a deliberate edit to this line rather than something that slips in with a lockfile. + const ALLOWED = new Set(['onnxruntime-web']); const pkg = JSON.parse(readFileSync(join(REPO, 'package.json'), 'utf8')); const deps = Object.keys({ ...pkg.dependencies, ...pkg.devDependencies }); - expect(deps.filter((d) => /onnx|tensorflow|tfjs|torch|tflite/i.test(d))).toEqual([]); - for (const f of ['xgbEnsemble.js', 'hitLikelihoodModel.js', 'hitLikelihood.js', 'scope.js']) { - const src = readFileSync(join(PRESCREEN, f), 'utf8'); - expect(src, `${f} references an ML runtime`).not.toMatch( - /\bonnx|onnxruntime|ort-wasm|tensorflow|tfjs\b/i - ); - } + const unexpected = deps.filter((d) => RUNTIME.test(d) && !ALLOWED.has(d)); + expect(unexpected, 'unexpected ML runtime in package.json').toEqual([]); }); }); @@ -1633,7 +1730,10 @@ describe.skipIf(!oos)('out-of-sample calibration (the rates the copy quotes)', ( console.log( `[meme-hit-likelihood] OUT-OF-SAMPLE on ${oos.length} unseen jobs: ` + ['unlikely', 'uncertain', 'likely'] - .map((b) => `${b} ${band[b].hit}/${band[b].n} = ${(100 * band[b].hit / band[b].n).toFixed(1)}%`) + .map( + (b) => + `${b} ${band[b].hit}/${band[b].n} = ${((100 * band[b].hit) / band[b].n).toFixed(1)}%` + ) .join(' | ') ); }); @@ -1678,8 +1778,10 @@ describe.skipIf(!oos)('out-of-sample calibration (the rates the copy quotes)', ( const level = byValue[value]; if (!level) continue; // the base-rate claim is checked below const rate = band[level].hit / band[level].n; - expect(Math.abs(rate - value), `panel's ${what} (${value}) vs measured ${rate.toFixed(3)}`) - .toBeLessThan(0.06); + expect( + Math.abs(rate - value), + `panel's ${what} (${value}) vs measured ${rate.toFixed(3)}` + ).toBeLessThan(0.06); } }); @@ -1688,4 +1790,3 @@ describe.skipIf(!oos)('out-of-sample calibration (the rates the copy quotes)', ( expect(Math.abs(hits / oos.length - 0.85)).toBeLessThan(0.03); }); }); - diff --git a/src/test/tree-sanitation.test.js b/src/test/tree-sanitation.test.js new file mode 100644 index 0000000..8e23443 --- /dev/null +++ b/src/test/tree-sanitation.test.js @@ -0,0 +1,120 @@ +/** + * Tests for treeSanitation.js. + * + * The fixtures are not invented. They are the shapes DM3's own NJ inference produces + * (src/data/shared/NJ.bf) and the values measured to crash the AxoMEME 2.0 inference path at + * predict_regression_nexus.py:955, which computes log((node_count + 1.0) / (dist + 0.1)). + */ +import { describe, it, expect } from 'vitest'; +import { + branchLengths, + inspectBranchLengths, + hasCrashingBranchLength, + NJ_SATURATION_SENTINEL +} from '../lib/utils/treeSanitation.js'; + +describe('branchLengths', () => { + it('reads plain, scientific and negative lengths', () => { + expect(branchLengths('((a:0.1,b:0.2):0.05,c:0.3);')).toEqual([0.1, 0.2, 0.05, 0.3]); + expect(branchLengths('(a:1.5e-2,b:3.0E-3);')).toEqual([0.015, 0.003]); + expect(branchLengths('(x:-0.5,y:0.2,z:0.4);')).toEqual([-0.5, 0.2, 0.4]); + }); + + it('returns nothing for a topology-only tree or non-string input', () => { + expect(branchLengths('((a,b),c);')).toEqual([]); + expect(branchLengths('')).toEqual([]); + expect(branchLengths(null)).toEqual([]); + }); + + it('is not confused by a bootstrap value or a quoted label', () => { + // )95: is a support value on an internal node, not a length; 'Homo:sapiens' is a name. + expect(branchLengths('((a:0.1,b:0.2)95:0.05,c:0.3);')).toEqual([0.1, 0.2, 0.05, 0.3]); + expect(branchLengths("(('Homo:sapiens':0.1,b:0.2):0.05);")).toEqual([0.1, 0.2, 0.05]); + }); + + it('is re-entrant — a global regex must not carry lastIndex between calls', () => { + const t = '((a:0.1,b:0.2):0.05,c:0.3);'; + expect(branchLengths(t)).toEqual(branchLengths(t)); + }); +}); + +describe('inspectBranchLengths', () => { + it('passes a clean tree', () => { + const r = inspectBranchLengths('((a:0.1,b:0.2):0.05,c:0.3);'); + expect(r.ok).toBe(true); + expect(r.negative).toBe(0); + expect(r.hasLengths).toBe(true); + expect(r.reasons).toEqual([]); + }); + + it('flags a topology-only tree', () => { + const r = inspectBranchLengths('((a,b),c);'); + expect(r.ok).toBe(false); + expect(r.hasLengths).toBe(false); + expect(r.reasons[0]).toMatch(/topology-only/); + }); + + it('counts negatives and reports the worst one', () => { + const r = inspectBranchLengths('(x:-0.5,y:0.2,z:-0.01);'); + expect(r.negative).toBe(2); + expect(r.total).toBe(3); + expect(r.min).toBe(-0.5); + expect(r.negativeFraction).toBeCloseTo(2 / 3, 10); + expect(r.reasons.join(' ')).toMatch(/negative/); + }); + + it('flags the NJ saturation sentinel, which is not a distance', () => { + // NJ.bf:99 returns 1000 for a saturated pair; the three-taxon closed form at NJ.bf:214-220 + // turns that into roughly -499.9. + const r = inspectBranchLengths(`(a:${NJ_SATURATION_SENTINEL},b:0.2);`); + expect(r.saturated).toBe(1); + expect(r.ok).toBe(false); + expect(r.reasons.join(' ')).toMatch(/saturation sentinel/); + + const derived = inspectBranchLengths('(a:-499.9,b:0.2);'); + expect(derived.negative).toBe(1); + expect(derived.ok).toBe(false); + }); + + it('does not modify the tree it inspects', () => { + // The module reports; it must never silently rewrite branch lengths, because those are the + // input a consumer's distances are built from. + const t = '(x:-0.5,y:0.2);'; + inspectBranchLengths(t); + expect(t).toBe('(x:-0.5,y:0.2);'); + }); +}); + +describe('hasCrashingBranchLength', () => { + // Measured against the real inference path: log((3 + 1.0) / (d + 0.1)) is a math domain error + // for d <= -0.1. -0.010 survives; -0.11, -0.5 and -499.9 all crash. + it('matches the measured crash threshold', () => { + expect(hasCrashingBranchLength('(a:-0.010,b:0.2);')).toBe(false); + expect(hasCrashingBranchLength('(a:-0.11,b:0.2);')).toBe(true); + expect(hasCrashingBranchLength('(a:-0.5,b:0.2);')).toBe(true); + expect(hasCrashingBranchLength('(a:-499.9,b:0.2);')).toBe(true); + }); + + it('is false for clean and topology-only trees', () => { + expect(hasCrashingBranchLength('((a:0.1,b:0.2):0.05,c:0.3);')).toBe(false); + expect(hasCrashingBranchLength('((a,b),c);')).toBe(false); + }); + + it('honours a different epsilon', () => { + expect(hasCrashingBranchLength('(a:-0.3,b:0.2);', 0.5)).toBe(false); + expect(hasCrashingBranchLength('(a:-0.6,b:0.2);', 0.5)).toBe(true); + }); + + it('is documented as necessary but NOT sufficient', () => { + // The consumer's threshold is on PATRISTIC distances — path sums — which can be more + // negative than any single branch. Two small negatives on one path sum past the threshold + // while no individual branch does. This test pins that limitation so nobody reads a false + // return as a safety guarantee. + const t = '((a:-0.06,b:-0.06):0.01,c:0.3);'; + expect(hasCrashingBranchLength(t)).toBe(false); // no single branch reaches -0.1 + const sum = branchLengths(t) + .filter((v) => v < 0) + .reduce((s, v) => s + v, 0); + expect(sum).toBeLessThan(-0.1); // but the path they share does + }); +}); diff --git a/static/models/axomeme/axomeme_2.0_viral_finetuned.onnx b/static/models/axomeme/axomeme_2.0_viral_finetuned.onnx new file mode 100644 index 0000000..c3f554a Binary files /dev/null and b/static/models/axomeme/axomeme_2.0_viral_finetuned.onnx differ diff --git a/vite.config.ts b/vite.config.ts index b399928..99f27de 100644 --- a/vite.config.ts +++ b/vite.config.ts @@ -67,6 +67,12 @@ export default defineConfig({ optimizeDeps: { include: ['@biowasm/aioli', 'toml', 'marked', 'socket.io-client'], // Exclude linked packages so changes are picked up immediately. + // + // hyphy-scope is deliberately NOT in this list even when linked via `npm link`. It pulls d3, + // phylotree, circos and @observablehq/plot, and excluding it from pre-bundling makes the dev + // server transform all of that on every request — enough to stall the page, which presents as + // an analysis that never finishes rather than as a slow import. Pre-bundle it and pick up + // library changes by rebuilding hyphy-scope and restarting dev (or `vite dev --force`). exclude: ['phylotree', 'alivibe'] } });