@wlearn/mitra
v0.3.0
Published
Mitra Tab2D in-context learning model as a wlearn Estimator (ONNX Runtime)
Maintainers
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
.wlrnbundle when you callsave() - 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 # BrowserYou 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.txtConvert
python convert.py # both variants
python convert.py --variant classifier
python convert.py --variant regressorProduces 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-4ONNX 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 testTests 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.
