Representax
Representax is a native JAX and Equinox system for efficient, task-general representation learning. Retrieval is the first working task; classification, reward modeling, distillation, and self-supervised objectives are planned on the same core boundary.
The project is alpha. The current slice provides:
- an Equinox-native encoder protocol with typed routes;
- a native ModernVBERT text-image encoder with bidirectional Hugging Face weight maps for every tensor used by its forward pass;
- direct multiple-negatives ranking, including symmetric and Matryoshka modes;
- an end-to-end Grain-to-compiled-step trainer with asynchronous reporting;
- lazy Grain recipes with built-in Hugging Face and local artifact resolvers;
- validated domain configs with annotated scientific and execution parameters; and
- explicit unit, runtime, parity, distributed, and performance test lanes.
Principles
- Native Equinox models are the supported execution path.
- Upstream implementations are optional development-time parity oracles.
- Scientific intent is separate from topology-dependent execution choices.
- Data recipes point at immutable artifacts and map them lazily into task examples; they do not require proprietary materialized datasets.
- Text, image, audio, video, and fused inputs must be supported without changing the training abstractions.
Install
The base package installs the ordinary CPU JAX runtime, Grain data pipeline, and safetensor checkpoint support. It never installs PyTorch:
python -m pip install representax
python -m pip install "representax[cuda13]" # NVIDIA GPU, recommended
python -m pip install "representax[cuda12]" # older NVIDIA drivers
For an editable source checkout, put -e . in place of representax.
Optional capabilities are grouped deliberately. Upstream parity oracles are
repository-only dependency groups rather than published package extras:
python -m pip install -e ".[config,hf]"
python -m pip install -e ".[test]" --group parity
python -m pip install -e ".[test]" --group parity-modernvbert
python -m pip install -e ".[test,performance]" --group parity-modernvbert
The v0 Hugging Face reference is pinned to Transformers 5.3.0. Its complete architecture catalog is distinct from native support: ModernVBERT is currently the only architecture carrying the verified native-support claim.
See the compatibility matrix for the locally accepted Python/JAX combinations and the distinction between CPU CI and accelerator acceptance. Maintainers should follow the release procedure for artifact inspection and trusted publication.
Gotchas
Pip-managed CUDA can be shadowed by LD_LIBRARY_PATH
The cuda12 and cuda13 extras use JAX's pip-managed NVIDIA libraries. A
shell-level LD_LIBRARY_PATH that points at another CUDA installation can take
precedence.
Did you see this?
Jax plugin configuration error: Exception when calling jax_plugins.xla_cuda12.initialize()
RuntimeError: Unable to load cuSPARSE. Is it installed?
Or did jax.devices() return only CpuDevice despite a working NVIDIA driver?
Then compare the ordinary process with one that ignores LD_LIBRARY_PATH:
python -c 'import jax; print(jax.devices())'
env -u LD_LIBRARY_PATH python -c 'import jax; print(jax.devices())'
If the second command restores the GPU, launch Representax the same way or unset the variable in that environment:
env -u LD_LIBRARY_PATH python train.py
# Or, for the current shell:
unset LD_LIBRARY_PATH
Do not remove a machine-wide CUDA configuration blindly: JAX's *-local
installations intentionally use a system toolkit. This advice applies to the
pip-managed extras documented above and follows JAX's
NVIDIA installation guidance.
Encoding
The compiled primitive has one route-aware operation:
import jax
import jax.numpy as jnp
import representax as rpx
model = rpx.models.DenseEncoder(8, 4, key=jax.random.key(0))
batch = jnp.ones((2, 8))
embeddings = rpx.encode(model, batch, route=rpx.Route.QUERY)
encode_documents = rpx.bind(model, route=rpx.Route.DOCUMENT)
document_embeddings = encode_documents(batch)
Host-side tokenization, media decoding, and batching will be exposed through a
higher-level embed API as production model integrations land.
ModernVBERT
The first production-family integration loads pinned Hugging Face safetensors directly into native Equinox text and SigLIP vision towers:
import jax.numpy as jnp
import representax as rpx
adapter = rpx.models.ModernVBERTCheckpointAdapter()
model = adapter.load("/path/to/modernvbert-checkpoint")
image_tokens = jnp.full((1, 64), 50407, dtype=jnp.int32)
input_ids = jnp.concatenate(
(jnp.asarray([[1]]), image_tokens, jnp.asarray([[2]])), axis=1
)
batch = rpx.models.ModernVBERTBatch(
input_ids=input_ids,
attention_mask=jnp.ones_like(input_ids),
pixel_values=jnp.ones((1, 1, 3, 512, 512), dtype=jnp.float32),
)
embeddings = rpx.encode(model, batch, route=rpx.Route.QUERY)
The real checkpoint uses 64 image tokens per processed 512x512 image. The
ordinary runtime includes safetensors but not PyTorch. A repository-only
pinned Transformers environment verifies vision features, fused
representations, and pixel gradients. Host-side Idefics3-compatible processing
remains the next API slice.
Versioned data recipes
Recipes are ordinary Python values that can be composed in Hydra-Zen config files and reviewed in Git:
from representax import data
recipe = data.mix(
data.source(
"hf://organization/dataset",
revision="immutable-revision",
split="train",
map="my_project.mappers.to_retrieval_example",
),
data.source(
"file:///data/corpus/train.parquet",
map="my_project.mappers.to_retrieval_example",
),
weights=(0.7, 0.3),
seed=17,
)
dataset = data.build_grain_dataset(recipe)
The recipe records artifact identity, mapping code identity, and sampling policy. A training iterator additionally fingerprints the resolved mapper and resolver implementations, batch mapper, batching contract, and Grain version. Grain performs lazy mapping, deterministic mixing, shuffling, and checkpointable iteration. A single source is the one-element form of the same sampling policy. Built-in resolvers support revision-pinned Hugging Face splits and local JSONL, Parquet, Arrow, or dataset directories. See the data contract for cache and extension behavior.
Training
Application code imports concrete operations from their owning modules:
import optax
from representax.config import (
BatchConfig,
CheckpointConfig,
ComponentConfig,
JobConfig,
LoggingConfig,
ModelConfig,
OptimizationConfig,
TrainingConfig,
)
from representax.data import build_grain_iterator
from representax.tasks import build_task
from representax.tasks.retrieval import MNRConfig, RetrievalConfig
from representax.train import build_train_step, make_train_state, run_training
optimizer = optax.adamw(learning_rate=1e-3)
job = JobConfig(
name="example",
model=ModelConfig(target="my_project.Model"),
task=RetrievalConfig(),
loss=MNRConfig(scale=20.0, symmetric=True, negative_scope="global"),
optimization=OptimizationConfig(
optimizer=ComponentConfig(
target="optax.adamw",
parameters={"learning_rate": 1e-3},
),
),
data=recipe,
training=TrainingConfig(
global_batch_size=32,
max_steps=10_000,
seed=17,
batch=BatchConfig(micro_batch_size=32),
),
logging=LoggingConfig(console_every=100),
checkpointing=CheckpointConfig(every=1_000, keep=3),
)
task = build_task(job.task, job.loss)
state = make_train_state(model, optimizer)
batches = build_grain_iterator(job.data, batch_size=32, batch_fn=collate)
result = run_training(
state=state,
step=build_train_step(task, optimizer),
batches=batches,
job=job,
run_directory="runs/example",
)
The loop records W&B-ready metric names such as train/loss, valid/loss, and
perf/... in metrics.jsonl, lifecycle events in events.jsonl, and final
status in run.json. A bounded reporter worker performs the device-to-host
metric transfer and fans the same ordered rows out to optional consumers without
placing a synchronization barrier in every training iteration. Checkpoints are
written by Orbax with at most one asynchronous save in flight. Recreate the same
recipe, model/state template, task/optimizer program, and batch source and pass
resume=True to continue from the latest complete checkpoint. See
the training contract.
Tests
pytest
pytest -m runtime
pytest -m parity
pytest -m distributed
pytest -m performance
Tests live outside the package and mirror its model, task, data, and runtime structure. Pytest markers select orthogonal runtime, parity, distributed, and performance lanes. The default command runs fast, dependency-light tests.
Performance acceptance is evaluated against a matched upstream implementation on pinned hardware. Compile time, steady-state work, and peak device memory are measured separately; see the test contract.
Roadmap
todo.org is the
canonical project roadmap and shared source of
truth. It tracks the production encoder port, parity gates, GradCache,
distributed training, checkpoint/resume, task-native audio/video, reward
modeling, JEPA, Profilax, and the systems-then-model research program.
Metadata
Release files for representax 0.0.1
For a detailed explanation of source distributions (sdists) and built distributions (wheels), please see the package formats documentation.
Source distribution (sdist)
| File | Size | Uploaded | |
|---|---|---|---|
| representax-0.0.1.tar.gz | 122.2 kB | Details |
Built distribution (wheel)
| File | Interpreter | ABI | Platform | Reset |
|---|---|---|---|---|
| representax-0.0.1-py3-none-any.whl | Python 3 | none | any | Details |
Total release size: 203.8 kB
Release files / representax-0.0.1.tar.gz
| Download URL | representax-0.0.1.tar.gz |
|---|---|
| Size | 122.2 kB |
| Tags | Source |
|
SHA-256 checksum How to use checksums |
56f68c804a734bb2dd41f70d24d8ef966dfc13c155fb612d438408237d6b0efe
|
|
BLAKE2b-256 checksum How to use checksums |
742f18fd616f490b5b67f5a2fe421d96a385421ffe9322b657b6c29147ddf2f4
|
| Upload date | |
|
Uploaded using Trusted Publishing? What is trusted publishing? |
Yes |
| Uploaded via |
twine/7.0.0 CPython/3.13.14
|
Provenance
Provenance describes where a file came from. On PyPI, provenance is shared via attestations, which provide a verifiable record of the build or publishing details. View details, limitations and caveats.
PyPI Publish Attestation
PyPI verified that this artifact, at this checksum, originated from the publisher listed below.
Signed by GitHub Actions, verified by PyPI on Aug 15, 2026.
Transparency logRelease files / representax-0.0.1-py3-none-any.whl
| Download URL | representax-0.0.1-py3-none-any.whl |
|---|---|
| Size | 81.6 kB |
| Tags | Python 3 |
|
SHA-256 checksum How to use checksums |
f51a7cc8cc81044897bb14b9a780c36ce7219ca22fa4f16412da8767b78b930b
|
|
BLAKE2b-256 checksum How to use checksums |
424dde141a31de175dee567d067a340651148a802857e4d87143800f54ae5494
|
| Upload date | |
|
Uploaded using Trusted Publishing? What is trusted publishing? |
Yes |
| Uploaded via |
twine/7.0.0 CPython/3.13.14
|
Provenance
Provenance describes where a file came from. On PyPI, provenance is shared via attestations, which provide a verifiable record of the build or publishing details. View details, limitations and caveats.
PyPI Publish Attestation
PyPI verified that this artifact, at this checksum, originated from the publisher listed below.
Signed by GitHub Actions, verified by PyPI on Aug 15, 2026.
Transparency log