Skip to main content

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

  1. Native Equinox models are the supported execution path.
  2. Upstream implementations are optional development-time parity oracles.
  3. Scientific intent is separate from topology-dependent execution choices.
  4. Data recipes point at immutable artifacts and map them lazily into task examples; they do not require proprietary materialized datasets.
  5. 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)

Source distribution for representax 0.0.1
File Size Uploaded
representax-0.0.1.tar.gz 122.2 kB Details

Built distribution (wheel)

Table of built distributions (wheels) for representax 0.0.1
File Interpreter ABI Platform
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 log

Release 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

Release history Release notifications | RSS feed

0.0.2

2 release files

This release

0.0.1 This release

2 release files

Anthropic, PBC Visionary sponsor Bloomberg Visionary sponsor Hudson River Trading Visionary sponsor Meta Visionary sponsor NVIDIA Visionary sponsor Microsoft Sustainability sponsor Depot Continuous Integration AWS Cloud computing and Security Sponsor Datadog Monitoring Fastly CDN Google Download Analytics Sentry Error logging StatusPage Status page