Skip to main content

A zero-RAM, multi-modal, sharded binary dataset manager

Project description

Roxxel 🚀

Zero-RAM JAX Dataloader, Asynchronous Checkpointer & Logging Trainer for Flax NNX

Roxxel is a zero-RAM, ultra-lightweight, and exceptionally fast dataset manager and pre-training orchestrator designed specifically for distributed JAX/Flax NNX clusters.

By utilizing virtualized POSIX memory-mapped dataset sharding, background asynchronous Orbax checkpointing, and thread-safe distributed logging, Roxxel completely does away with heavy, complex, and over-engineered training frameworks.


🌟 Key Features

  • OS-Level Memory Mapping (mmap): Maps multi-terabyte datasets directly into virtual memory via the operating system's kernel page cache. Consumes exactly 0 bytes of Python RAM for dataset storage.
  • Unified Causal Streaming: Automatically chunks, shuffles, and loads batches directly onto JAX device layouts (jax.device_put) using your Named Sharding mesh. Exposes the exact step count (len(stream)).
  • Instant Offset Seeking: Resumes streaming from any step index in under 1 millisecond using binary offsets—completely skipping the need to execute slow dummy fast-forward loops.
  • Dynamic Dtype Auto-Detection: Detects datatypes (e.g. int32 token IDs, float32 arrays, or raw text bytes) on compilation, writes them in a backward-compatible format, and decodes them perfectly on read.
  • Multi-Dataset Blending / Mixing: Blends a primary dataset with multiple secondary datasets using weight ratios. Automatically handles dataset exhaustion mid-phase by re-normalizing weights on the fly.
  • Curriculum Schedule Timelines: Supports dynamic sequence length extension (e.g., shifting from 1K to 32K context windows) and batch size changes at precise training step boundaries.
  • Topology-Agnostic Checkpointing: Offloads PyTree serialization asynchronously to background threads via Orbax Checkpoint Manager. Restores parameters natively using abstract templates, meaning model definition changes do not break older saved weights.
  • Multi-Host TPU/GPU Pod Safety: Restrictions ensure only Rank 0 writes stdout prints, system logs, and metrics CSV files, avoiding multi-process locking contention and terminal clutter.

🤖 What Roxxel Handles Automatically (So You Don't Have To)

Roxxel was engineered to remove the friction of writing accelerator-optimized JAX code. It handles the following tasks automatically under the hood:

  1. JIT Train Step Compilation: On initialization, Trainer automatically compiles your JIT training step (@nnx.jit). You don't need to write or trace JAX functions manually.
  2. Loss Unpacking & Wrapping: If your loss function returns a tuple/list (like (loss, aux_data)) or a dictionary (like {"loss": loss, "perplexity": ppl}), Roxxel automatically extracts the scalar loss for gradients (nnx.value_and_grad), avoiding compiler crashes.
  3. Auto-Initialized ModelState: You don't have to write state wrapper boilerplate (e.g. TrainState classes). The trainer dynamically bundles your model, optimizer, and training step counter.
  4. Auto-Initialized Checkpointer & Logger: If you pass a folder path to the save_path parameter, Roxxel automatically initializes the asynchronous Orbax Checkpointer and multi-threaded Logger under that directory, saving checkpoints in save_path/checkpoints and logs directly in save_path.
  5. Dynamic Stream Re-instantiation: During curriculum phase transitions (e.g. when sequence length changes), Roxxel automatically closes the active stream, computes completed offsets, and swaps the dataset streams instantly.
  6. Exhaustion Re-normalization: If one of your mixed datasets runs out of records mid-training, Roxxel removes it from the choice pool and re-normalizes the weights of the remaining active datasets to prevent training stalls.
  7. Crash-Safe Log Flushing: In the event of an OOM, crash, or exception, the logger intercepts the exception, logs the stack trace to the system log, and flushes all pending file and stdout writes before bubbling the error.

📦 Installation

To install the core dataloader and async logging engine only (no JAX required):

pip install roxxel

To install the JAX-native asynchronous trainer and checkpointing extensions:

pip install roxxel[checkpoint]

🚀 Quick Start Cookbook

Here is a complete, zero-boilerplate example showing how to initialize a model, define a curriculum schedule, and execute pre-training using the consolidated Trainer:

import jax
import optax
from flax import nnx
from roxxel import Roxxel, Phase, Curriculum, Trainer

# 1. Initialize Flax NNX model and optimizer
model = nnx.Linear(10, 5, rngs=nnx.Rngs(42))
tx = optax.sgd(0.01)
optimizer = nnx.Optimizer(model, tx, wrt=nnx.Param)

# 2. Define the curriculum (e.g., Phase 1: 1000 steps, Phase 2: 500 steps)
phases = [
    Phase(steps=1000, batch_size=16, seq_len=128),
    Phase(steps=500, batch_size=4, seq_len=512)
]
curriculum = Curriculum(
    primary_streamer=Roxxel("./wiki_tokens_*.rox"), 
    phases=phases
)

# 3. Define the loss function
def loss_fn(model, batch):
    logits = model(batch[:, :-1].astype(jax.numpy.float32))
    targets = batch[:, 1:].astype(jax.numpy.float32)
    return jax.numpy.mean((logits - targets) ** 2)

# 4. Initialize the Trainer
# Setting save_path automatically initializes Checkpointer, Logger, and ModelState
trainer = Trainer(
    model=model,
    optimizer=optimizer,
    curriculum=curriculum,
    loss_fn=loss_fn,
    save_path="./run_delta",
    checkpoint_every=100,
    log_every=10
)

# 5. Run curriculum training
trainer.run()

📖 Learn More

For complete documentation, design guides, and API specs, visit the Roxxel Documentation Site:


⚖️ License

MIT License. Feel free to use, modify, and distribute.

Project details


Download files

Download the file for your platform. If you're not sure which to choose, learn more about installing packages.

Source Distribution

roxxel-0.7.6.tar.gz (245.5 kB view details)

Uploaded Source

Built Distribution

If you're not sure about the file name format, learn more about wheel file names.

roxxel-0.7.6-py3-none-any.whl (22.3 kB view details)

Uploaded Python 3

File details

Details for the file roxxel-0.7.6.tar.gz.

File metadata

  • Download URL: roxxel-0.7.6.tar.gz
  • Upload date:
  • Size: 245.5 kB
  • Tags: Source
  • Uploaded using Trusted Publishing? Yes
  • Uploaded via: twine/6.1.0 CPython/3.13.12

File hashes

Hashes for roxxel-0.7.6.tar.gz
Algorithm Hash digest
SHA256 c24f5361ca2c842748310f9fd7655d66caba2348491d0670a345ef65b9755d1d
MD5 711793f53e580626667abbd459bf9f68
BLAKE2b-256 35a0c125b83fd410ab16553b89b22bb3e7d7e9fb2671f4c884f291e1b25113ef

See more details on using hashes here.

Provenance

The following attestation bundles were made for roxxel-0.7.6.tar.gz:

Publisher: publish.yml on anon160/Roxxel

Attestations: Values shown here reflect the state when the release was signed and may no longer be current.

File details

Details for the file roxxel-0.7.6-py3-none-any.whl.

File metadata

  • Download URL: roxxel-0.7.6-py3-none-any.whl
  • Upload date:
  • Size: 22.3 kB
  • Tags: Python 3
  • Uploaded using Trusted Publishing? Yes
  • Uploaded via: twine/6.1.0 CPython/3.13.12

File hashes

Hashes for roxxel-0.7.6-py3-none-any.whl
Algorithm Hash digest
SHA256 b7ca864b07c379972ef53bd306dd8b875ec22b7e1ea7b3b2d73d0e24f7de2d82
MD5 60039b589e933a882294146bde359abe
BLAKE2b-256 cfac2e2a64594fafd5de27e25d6cb90a409163f74f06ec43cc6316a8e05cdbb3

See more details on using hashes here.

Provenance

The following attestation bundles were made for roxxel-0.7.6-py3-none-any.whl:

Publisher: publish.yml on anon160/Roxxel

Attestations: Values shown here reflect the state when the release was signed and may no longer be current.

Supported by

AWS Cloud computing and Security Sponsor Datadog Monitoring Depot Continuous Integration Fastly CDN Google Download Analytics Pingdom Monitoring Sentry Error logging StatusPage Status page