npm package discovery and stats viewer.

Discover Tips

  • General search

    [free text search, go nuts!]

  • Package details

    pkg:[package-name]

  • User packages

    @[username]

Sponsor

Optimize Toolset

I’ve always been into building performant and accessible sites, but lately I’ve been taking it extremely seriously. So much so that I’ve been building a tool to help me optimize and monitor the sites that I build to make sure that I’m making an attempt to offer the best experience to those who visit them. If you’re into performant, accessible and SEO friendly sites, you might like it too! You can check it out at Optimize Toolset.

About

Hi, 👋, I’m Ryan Hefner  and I built this site for me, and you! The goal of this site was to provide an easy way for me to check the stats on my npm packages, both for prioritizing issues and updates, and to give me a little kick in the pants to keep up on stuff.

As I was building it, I realized that I was actually using the tool to build the tool, and figured I might as well put this out there and hopefully others will find it to be a fast and useful way to search and browse npm packages as I have.

If you’re interested in other things I’m working on, follow me on Twitter or check out the open source projects I’ve been publishing on GitHub.

I am also working on a Twitter bot for this site to tweet the most popular, newest, random packages from npm. Please follow that account now and it will start sending out packages soon–ish.

Open Software & Tools

This site wouldn’t be possible without the immense generosity and tireless efforts from the people who make contributions to the world and share their work via open source initiatives. Thank you 🙏

© 2026 – Pkg Stats / Ryan Hefner

@wlearn/mitra

v0.3.0

Published

Mitra Tab2D in-context learning model as a wlearn Estimator (ONNX Runtime)

Readme

@wlearn/mitra

Mitra Tab2D in-context learning model as a wlearn (GitHub, all packages) Estimator. Wraps the 72M-parameter tabular foundation model from Amazon/AutoGluon in the standard wlearn fit() / predict() / save() / load() API.

Mitra uses in-context learning: instead of gradient-based training, fit() selects a support set from your data, and predict() passes that support set alongside your query data through the ONNX model in a single forward pass.

Data storage warning

fit() stores a subset of your training data (up to maxSupport rows, default 512) inside the model instance. This support set is:

  • Held in memory as Float32Array for the lifetime of the instance
  • Serialized into the .wlrn bundle when you call save()
  • Required for every predict() call (it is the model's "context")

If your training data is sensitive, be aware that saved bundles contain real data points. The maxSupport parameter controls how many rows are stored. Use dispose() when you need to release the support set before the object is collected.

Inference memory

The ONNX forward pass can use substantially more memory than the support arrays or saved bundle. In a bounded Linux probe with Node 22, onnxruntime-node 1.24.2, the 303 MB classifier, and 64-feature Digits data, three repeated predictions reached kernel peak RSS of about 1.02 GiB for 50 support / 10 query rows and 1.48 GiB for 100 support / 20 query rows. The third prediction did not increase RSS beyond the second in either case, consistent with runtime arena warm-up, but these measurements are environment-specific rather than a guaranteed bound.

Start with a small maxSupport and query size, then measure the intended runtime. This wrapper does not silently split query batches because numerical equivalence for this model has not been established. dispose() clears the wlearn-owned support set. If you supplied an InferenceSession, release it yourself after all models using it are disposed; that session owns the ONNX Runtime allocation arena. For a byte-backed model, dispose() clears the support data synchronously and initiates the owned session's asynchronous ONNX Runtime release. The method remains synchronous; any release rejection is observed and reported as a warning.

Install

npm install @wlearn/mitra onnxruntime-node   # Node.js
npm install @wlearn/mitra onnxruntime-web    # Browser

You also need the ONNX model files (see "ONNX conversion" below).

Usage

Classifier

Mitra needs an ONNX source during construction and loading, so use the explicit MitraClassifier or MitraRegressor class. There is no generic factory because the classifier and regressor require different ONNX sources.

const { MitraClassifier } = require('@wlearn/mitra')
const ort = require('onnxruntime-node')
const { readFileSync, writeFileSync } = require('node:fs')

async function main() {
  const XTrain = [[0, 0], [0, 1], [1, 0], [1, 1]]
  const yTrain = new Int32Array([0, 0, 1, 1])
  const XTest = [[0.25, 0.25], [0.75, 0.75]]
  const yTest = new Int32Array([0, 1])

  // The ONNX network is distributed separately from the npm package.
  const onnxBytes = readFileSync('mitra-classifier.onnx')
  const runtimeOptions = { ort, sessionOptions: { intraOpNumThreads: 2, interOpNumThreads: 1 } }
  const model = await MitraClassifier.create(onnxBytes, {
    maxSupport: 4,
    seed: 42
  }, runtimeOptions)

  model.fit(XTrain, yTrain) // synchronous support-set selection
  console.log(Array.from(await model.predict(XTest)))
  console.log(Array.from(await model.predictProba(XTest)))
  console.log(await model.score(XTest, yTest))

  const bundle = model.save()
  writeFileSync('classifier.wlrn', bundle)
  model.dispose()

  const loaded = await MitraClassifier.load(
    readFileSync('classifier.wlrn'), onnxBytes, runtimeOptions
  )
  console.log(Array.from(await loaded.predict(XTest)))

  loaded.dispose()
}

main().catch(error => {
  console.error(error)
  process.exitCode = 1
})

The WLRN bundle stores the selected support rows but not the ONNX network. Passing ONNX bytes lets Mitra hash the exact bytes and bind the bundle to that identity; load rejects different bytes before creating a session. Node Buffer values from readFileSync() are accepted for both ONNX and WLRN files. Runtime options accept sessionOptions, forwarded to ONNX Runtime during byte-backed construction and loading; use it to select execution providers or limit CPU threads.

A pre-created InferenceSession is also accepted, but its bytes are no longer available to inspect. To save or load an identity-bound bundle in that mode, pass the caller-verified model hash as { trustedOnnxSha256 }. The hash is an explicit caller assertion, not a value observed from the session.

Regressor

This focused helper is a standalone CommonJS fragment; supply your own data when calling it:

const { MitraRegressor } = require('@wlearn/mitra')
const ort = require('onnxruntime-node')
const { readFileSync } = require('node:fs')

async function runRegressor(XTrain, yTrain, XTest, yTest) {
  const onnxBytes = readFileSync('mitra-regressor.onnx')
  const model = await MitraRegressor.create(
    onnxBytes, { maxSupport: 50 }, { ort, sessionOptions: { intraOpNumThreads: 2, interOpNumThreads: 1 } }
  )
  try {
    model.fit(XTrain, yTrain)
    return {
      predictions: await model.predict(XTest),
      r2: await model.score(XTest, yTest)
    }
  } finally {
    model.dispose()
  }
}

Registry integration

Register each available task with its matching ONNX bytes. Registry load() is asynchronous for Mitra because loading constructs an ONNX Runtime session:

const { registerLoaders } = require('@wlearn/mitra')
const { load } = require('@wlearn/core')

async function loadMitraBundle(bundleBytes, classifierOnnx, regressorOnnx, ort) {
  registerLoaders(classifierOnnx, regressorOnnx, { ort, sessionOptions: { intraOpNumThreads: 2, interOpNumThreads: 1 } })
  return await load(bundleBytes)
}

For pre-created sessions, pass the caller-verified task hashes as classifierTrustedOnnxSha256 and regressorTrustedOnnxSha256 in the final options object.

API

MitraClassifier

| Method | Returns | Description | |--------|---------|-------------| | static create(onnxSource, params?, opts?) | Promise<MitraClassifier> | Factory. onnxSource is an InferenceSession or Uint8Array of ONNX bytes | | fit(X, y) | this | Select support set from training data (sync) | | predict(X) | Promise<Int32Array> | Class label predictions | | predictProba(X) | Promise<Float64Array> | Class probabilities (rows * nClasses) | | score(X, y) | Promise<number> | Accuracy | | save() | Uint8Array | Serialize to .wlrn bundle | | static load(bytes, onnxSource, opts?) | Promise<MitraClassifier> | Deserialize | | dispose() | void | Clear support data; initiate session release only when Mitra created it from ONNX bytes | | getParams() | object | { maxSupport, seed } | | setParams(p) | this | Update params before fitting | | capabilities | object | { classifier: true, predictProba: true, ... } | | classes | Int32Array | Sorted unique class labels | | nrClass | number | Number of classes | | nrFeature | number | Number of features |

When construction receives a pre-created session, the caller retains ownership and must release that session after disposing every Mitra object that uses it. Saving a session-backed model requires opts.trustedOnnxSha256; byte-backed models record the observed SHA-256 automatically.

MitraRegressor

Same API minus predictProba, classes, nrClass. score() returns R2.

Parameters

| Param | Default | Description | |-------|---------|-------------| | maxSupport | 512 | Maximum support set size. If training data exceeds this, a subset is sampled | | seed | 42 | RNG seed for deterministic support set sampling |

Bundle format

The .wlrn bundle stores the support set, not the much larger external ONNX model (the measured classifier file is 303 MB). Loading requires the same ONNX model to be provided separately.

| Artifact | Format | Contents | |----------|--------|----------| | meta | JSON | Support shape, classes, ordinal-label encoding, seed, and ONNX identity evidence | | support_x | raw float32 | Support set features, row-major | | support_y | raw int32/float32 | Support set labels (int32 for classifier, float32 for regressor) |

Current writers use wlearn.mitra_onnx.classifier@2 and wlearn.mitra_onnx.regressor@2. Readers also accept identity-bearing @1 bundles. A legacy @1 bundle with no ONNX hash requires the explicit allowUnverifiedLegacyModel: true load option.

Migration from 0.2

Version 0.3 uses the task-specific MitraClassifier and MitraRegressor classes. The former generic MitraModel entry point was removed because it could save an ambiguous artifact whose task identity was not enforced. Current @2 bundles bind both the estimator task and the supplied ONNX model identity; legacy @1 loading follows the explicit policy above.

ONNX model variants

| Model | HuggingFace | Output | |-------|-------------|--------| | Classifier | autogluon/mitra-classifier | (B, N_query, 10) logits | | Regressor | autogluon/mitra-regressor | (B, N_query) values |

ONNX conversion

pip install -r requirements.txt

Convert

python convert.py              # both variants
python convert.py --variant classifier
python convert.py --variant regressor

Produces mitra-classifier.onnx and/or mitra-regressor.onnx.

Download published conversions

From this repository checkout, node scripts/download-models.mjs --dir DIR downloads the release named by models.json. It verifies existing files, streams the SHA-256 check, and installs only a complete verified download. A corrupted existing file is preserved and reported; move it before retrying. The shell entry point delegates to the same implementation. This is a repository command, not an npm executable.

Release assets must be uploaded during the coordinated release. If the requested asset returns HTTP 404, use the conversion commands above with the upstream weights. Keep locally converted files and verify them before redistribution.

Verify

Compare PyTorch and ONNX Runtime outputs:

python verify.py
python verify.py --atol 1e-4

ONNX model inputs

All dimensions are dynamic (variable batch, support/query/feature counts).

| Input | Shape | Type | |-------|-------|------| | x_support | (B, N_support, N_features) | float32 | | y_support | (B, N_support) | int64 (classifier) / float32 (regressor) | | x_query | (B, N_query, N_features) | float32 | | padding_obs_support | (B, N_support) | bool |

Note: padding_features and padding_obs_query are accepted by the PyTorch model for API compatibility but are unused in the CPU code path (they only matter for flash attention). The ONNX tracer correctly eliminates them from the graph.

Testing

Requires ONNX models in the repo root (run python convert.py first).

npm install
npm test

Tests cover create, fit, predict, predictProba, score, save/load round-trip, dispose, error handling, support set selection, and determinism.

What was changed for ONNX export

The upstream AutoGluon implementation (_internal/models/tab2d.py, _internal/models/embedding.py) uses several ops that the ONNX exporter cannot trace or has no opset mapping for. Each replacement below preserves numerical equivalence while using only standard ONNX ops.

Quantile computation (Tab2DQuantileEmbeddingX)

The upstream computes 999 quantiles of x_support along the observation axis with torch.quantile, which has no ONNX opset mapping.

Replacement: torch.sort along dim 1, then torch.gather at fractional index positions with linear interpolation between the floor and ceil indices. Produces identical quantile boundaries.

Bucketize / searchsorted (Tab2DQuantileEmbeddingX)

The upstream maps each value to its quantile bin using torch.vmap(torch.bucketize, in_dims=(0,0)). Both vmap (not traceable) and bucketize / searchsorted (no ONNX op) are unavailable.

Replacement: broadcasting comparison. For values (b, f, s) and boundaries (b, f, 999), compute (values.unsqueeze(-1) >= boundaries.unsqueeze(-2)).sum(-1). This counts how many boundaries each value exceeds, which is exactly the bucket index. O(n * 999) instead of O(n * log 999), but 999 is small and the operation is pure element-wise ONNX ops (GreaterOrEqual, Cast, ReduceSum).

In-place masked assignment (Tab2DQuantileEmbeddingX, Tab2DEmbeddingY*)

The upstream uses x_support[padding_mask] = 9999 and y_support[padding_obs] = 0 to set padded positions. In-place mutation through boolean indexing is not traceable.

Replacement: torch.where(mask, fill_value, x). Functionally identical, produces a new tensor instead of mutating.

einops / einx (Tab2D, Layer, MultiheadAttention, embeddings)

The upstream uses einops.rearrange, einops.pack/unpack, einx.rearrange, and einx.sum throughout. These are external libraries the ONNX tracer cannot see through.

Replacements:

  • einx.rearrange("b s f -> b s f 1", x) -- x.unsqueeze(-1)
  • einops.rearrange("b n -> b n 1", y) -- y.unsqueeze(-1)
  • einops.rearrange("b s f d -> (b f) s d", x) -- x.permute(0,2,1,3).reshape(b*f, s, d)
  • einops.rearrange("(b f) s d -> b s f d", x, b=b) -- x.reshape(b, f, s, d).permute(0,2,1,3)
  • einops.rearrange("b s f d -> (b s) f d", x) -- x.reshape(b*s, f, d)
  • einops.rearrange("b t (h d) -> b h t d", q, h=h) -- q.reshape(b, t, h, d).permute(0,2,1,3)
  • einops.pack((y, x), "b s * d") -- torch.cat([y, x], dim=2)
  • einops.unpack(q, pack_info, "b s * c") -- q[:, :, 0, :] (index the y slot)
  • einx.sum("b [s] f", x) -- x.sum(dim=1, keepdim=True)

Gradient checkpointing (Tab2D)

The upstream wraps each layer call in torch.utils.checkpoint.checkpoint(layer, ...) which is a training-only optimization not compatible with tracing.

Replacement: direct call layer(support, query). The model is exported in eval mode so checkpointing has no effect anyway.

Flash attention path (Tab2D, Layer, Padder)

The upstream has two code paths: a flash attention path (CUDA with flash_attn library) and a CPU path using F.scaled_dot_product_attention. The flash attention path uses flash_attn_varlen_func, unpad_input/pad_input, and a Padder class -- none of which are ONNX-exportable.

Replacement: only the CPU path is reimplemented. F.scaled_dot_product_attention maps cleanly to standard ONNX attention ops. The Padder class and all flash attention imports are removed entirely.

Summary table

| Original | Replacement | Location | |----------|-------------|----------| | torch.quantile | sort + gather + lerp | Tab2DQuantileEmbeddingX | | torch.vmap(torch.bucketize) | broadcast compare + sum | Tab2DQuantileEmbeddingX | | x[mask] = val (in-place) | torch.where | embeddings | | einops.rearrange / einx.rearrange | reshape / permute / unsqueeze | everywhere | | einops.pack / unpack | torch.cat / indexing | Tab2D.forward | | einx.sum | torch.sum | Tab2DQuantileEmbeddingX | | checkpoint(layer, ...) | layer(...) | Tab2D.forward | | Flash attention path + Padder | Removed (CPU path only) | Tab2D, Layer |

State dict keys match upstream exactly -- safetensors load without any key renaming.

License

The Mitra model weights are Apache-2.0 licensed by Amazon/AutoGluon. This conversion code is Apache-2.0 licensed.