dantinox.paradigms

Paradigms define the training objective and generation strategy. The Trainer only ever calls loss_fn — all paradigm-specific logic is self-contained.

Overview

A paradigm is a self-contained unit that owns:

  1. Model construction — build_model() returns the JAX/NNX model for this paradigm.

  2. Loss function — loss_fn(model, batch) returns (loss, metrics). The Trainer never touches the model directly — it calls this.

  3. Generation — generate() wraps the model-specific decode loop.

  4. Parameter count — num_parameters(model).

┌─────────────┐       loss_fn(model, batch)        ┌────────────┐
│   Trainer   │  ──────────────────────────────►   │  Paradigm  │
│             │  ◄── (loss: float, metrics: dict) ──│            │
└─────────────┘                                     └────────────┘

Unified Paradigm API

The recommended entry point is the single Paradigm class. Pass a ModelConfig with a paradigm key and the right implementation is selected automatically:

import dantinox as dx

# Autoregressive (causal=True set automatically)
p = dx.Paradigm(dx.ModelConfig(paradigm="ar", dim=512, n_heads=8, num_blocks=12))

# Discrete Diffusion (causal=False set automatically)
p = dx.Paradigm(dx.ModelConfig(paradigm="discrete", dim=512, n_heads=8, num_blocks=12,
                                noise_schedule="cosine", mask_token_id=4))

# Continuous Flow-Matching (causal=False set automatically)
p = dx.Paradigm(dx.ModelConfig(paradigm="continuous", dim=256, n_heads=4,
                                embed_dim=768, bottleneck_dim=128, num_blocks=6))

# Sentence Embedder
p = dx.Paradigm(dx.ModelConfig(paradigm="embedder", dim=512, n_heads=8, num_blocks=12,
                                embed_pooling="mean", embed_temperature=0.05))

Paradigm selection table

config.paradigm

Implementation

causal auto-set

Model class

"ar"

ARParadigm

True

Transformer (causal)

"discrete"

DiscreteParadigm

False

DiffusionTransformer

"continuous"

ContinuousParadigm

False

FlowMatchingTransformer

"embedder"

EmbedderParadigm

True

Transformer (pooled)

When paradigm=None (omitted), the implementation is auto-detected from causal and embed_dim for backward compatibility.


Paradigm reference

class dantinox.paradigms.paradigm.Paradigm(config: ModelConfig)[source]

Bases: ParadigmBase

Unified paradigm that routes to the right implementation via ModelConfig.paradigm.

Quick-start:

import dantinox as dx

# Autoregressive
p = dx.Paradigm(dx.ModelConfig(paradigm="ar", dim=512, n_heads=8, num_blocks=12))

# Discrete diffusion (LLaDA)
p = dx.Paradigm(dx.ModelConfig(paradigm="discrete", dim=512, n_heads=8, num_blocks=12,
                                noise_schedule="cosine"))

# Continuous flow-matching
p = dx.Paradigm(dx.ModelConfig(paradigm="continuous", dim=512, n_heads=8, num_blocks=12,
                                embed_dim=768))

# Contrastive embedder (SimCSE)
p = dx.Paradigm(dx.ModelConfig(paradigm="embedder", dim=512, n_heads=8, num_blocks=12,
                                dropout=0.1))

run_dir = dx.Trainer(p, dx.TrainingConfig(lr=3e-4, epochs=5)).fit("corpus.txt")

Setting paradigm in ModelConfig also auto-configures causal, so you never need to write causal=False for discrete/continuous.

Every concrete paradigm accepts its architecture config positionally.

Declared here (not abstract) so callers can do type(some_paradigm)(new_config) — e.g. the W&B sweep agent rebuilding a paradigm with overridden hyperparameters — without a static type error, even though this base implementation is never actually used (each subclass overrides it; some, like EmbedderParadigm, add further keyword-only arguments with defaults, so calling with just config remains valid everywhere).

build_model(rngs: Any) → Any[source]

Construct and return the NNX model for this paradigm.

Called once by the Trainer at the start of fit(). The returned model is then managed (checkpointed, sharded) by the Trainer and passed back into loss_fn / generate as the first argument.

loss_fn(model: Any, batch: Any, rng: Any, **kwargs: Any) → Any[source]

Compute the scalar training loss for one batch.

The model is passed explicitly so nnx.value_and_grad can differentiate through it without the paradigm needing to be an NNX module itself.

Parameters:

embeddings – Per-batch extras from prepare_batch, only passed when provides_batch_extras is True (e.g. pre-computed T5 embeddings for the continuous flow-matching paradigm). Ignored by paradigms that don’t declare provides_batch_extras = True.

Returns:

(scalar_loss, metrics_dict) where metrics_dict holds any auxiliary scalars (ce_loss, aux_loss, mse_loss, …) for logging.

generate(model: Any, *args: Any, **kwargs: Any) → Any[source]

Generate token sequences given a prompt prefix.

property provides_batch_extras: bool

bool(x) -> bool

Returns True when the argument x is true, False otherwise. The builtins True and False are the only two instances of the class bool. The class bool is a subclass of the class int, and cannot be subclassed.

property diffusion_config: Any

Diffusion-specific configuration. Returns None for non-diffusion paradigms.

property requires_shifted_targets: bool

bool(x) -> bool

Returns True when the argument x is true, False otherwise. The builtins True and False are the only two instances of the class bool. The class bool is a subclass of the class int, and cannot be subclassed.

on_train_start(model: Any, sample_batches: list) → None[source]

One-time setup before training starts (default: no-op).

Intentionally concrete, not abstract: most paradigms have nothing to do here. sample_batches is a small list of token batches drawn from the training set, for paradigms that need data-dependent initialisation.

prepare_batch(batch: Any) → Any[source]

Host-side per-batch preprocessing executed outside JIT.

Only called when provides_batch_extras is True; the return value is forwarded to loss_fn as the embeddings keyword argument.

stream(model: Any, *args: Any, **kwargs: Any) → Iterator[Any][source]

Yield (step, total, tokens) after each generation step.

Only available for "discrete" and "continuous" paradigms.

property type: str

"ar", "discrete", "continuous", or "embedder".

Type:

Return the active paradigm type


Base class

ParadigmBase is the abstract base class. Subclass it to implement a custom paradigm:

from dantinox.paradigms.base import ParadigmBase

class MyParadigm(ParadigmBase):
    def build_model(self, rngs):
        return MyModel(self.config, rngs=rngs)

    def loss_fn(self, model, batch, rng, **kwargs):
        logits = model(batch["input_ids"])
        loss   = cross_entropy(logits, batch["labels"])
        return loss, {"loss": loss}

    def generate(self, model, *args, **kwargs):
        return greedy_decode(model, *args, **kwargs)

    def num_parameters(self, model):
        return sum(x.size for x in jax.tree_util.tree_leaves(nnx.state(model, nnx.Param)))
class dantinox.paradigms.base.ParadigmBase(config: Any)[source]

Bases: ABC

Abstract base for all generative paradigms.

A Paradigm wraps a core model and owns the paradigm-specific logic: how to corrupt inputs, compute loss, and generate samples. The trainer calls only loss_fn and generate — it never inspects the internals.

Implementing a new paradigm requires overriding three methods:

class MyParadigm(Paradigm):
    def build_model(self, rngs): ...
    def loss_fn(self, model, batch, rng): ...
    def generate(self, model, prompt, rng, **kwargs): ...

Two optional hooks let a paradigm participate in the training loop without the Trainer knowing its internals:

  • on_train_start(model, sample_batches) — one-time setup before the first step (e.g. flow-matching computes T5 embedding normalisation statistics).

  • prepare_batch(batch) — host-side, non-JIT preprocessing of each batch; whatever it returns is forwarded to loss_fn as the embeddings keyword. Set provides_batch_extras = True to enable it.

Every concrete paradigm accepts its architecture config positionally.

Declared here (not abstract) so callers can do type(some_paradigm)(new_config) — e.g. the W&B sweep agent rebuilding a paradigm with overridden hyperparameters — without a static type error, even though this base implementation is never actually used (each subclass overrides it; some, like EmbedderParadigm, add further keyword-only arguments with defaults, so calling with just config remains valid everywhere).

requires_shifted_targets: bool = False
provides_batch_extras: bool = False
__init__(config: Any) → None[source]

Every concrete paradigm accepts its architecture config positionally.

Declared here (not abstract) so callers can do type(some_paradigm)(new_config) — e.g. the W&B sweep agent rebuilding a paradigm with overridden hyperparameters — without a static type error, even though this base implementation is never actually used (each subclass overrides it; some, like EmbedderParadigm, add further keyword-only arguments with defaults, so calling with just config remains valid everywhere).

abstractmethod build_model(rngs: Any) → Any[source]

Construct and return the NNX model for this paradigm.

Called once by the Trainer at the start of fit(). The returned model is then managed (checkpointed, sharded) by the Trainer and passed back into loss_fn / generate as the first argument.

abstractmethod loss_fn(model: Any, batch: Array, rng: Array, embeddings: Array | None = None) → tuple[Array, dict[str, Any]][source]

Compute the scalar training loss for one batch.

The model is passed explicitly so nnx.value_and_grad can differentiate through it without the paradigm needing to be an NNX module itself.

Parameters:

embeddings – Per-batch extras from prepare_batch, only passed when provides_batch_extras is True (e.g. pre-computed T5 embeddings for the continuous flow-matching paradigm). Ignored by paradigms that don’t declare provides_batch_extras = True.

Returns:

(scalar_loss, metrics_dict) where metrics_dict holds any auxiliary scalars (ce_loss, aux_loss, mse_loss, …) for logging.

abstractmethod generate(model: Any, prompt: Array, rng: Array, **kwargs: Any) → Array[source]

Generate token sequences given a prompt prefix.

on_train_start(model: Any, sample_batches: list[Any]) → None[source]

One-time setup before training starts (default: no-op).

Intentionally concrete, not abstract: most paradigms have nothing to do here. sample_batches is a small list of token batches drawn from the training set, for paradigms that need data-dependent initialisation.

prepare_batch(batch: Any) → Any[source]

Host-side per-batch preprocessing executed outside JIT.

Only called when provides_batch_extras is True; the return value is forwarded to loss_fn as the embeddings keyword argument.

num_parameters(model: Any) → int[source]

Count trainable parameters in the model.

property diffusion_config: Any

Diffusion-specific configuration. Returns None for non-diffusion paradigms.


Concrete implementations

The implementations below are directly importable for advanced use cases or when subclassing. In normal use, Paradigm wraps them transparently.

Autoregressive

class dantinox.paradigms.ar.ARParadigm(config: ModelConfig)[source]

Bases: ParadigmBase

Autoregressive next-token-prediction paradigm.

Loss: cross-entropy on shifted targets (teacher-forcing).

Quick-start:

cfg = ModelConfig(dim=512, n_heads=8, head_size=64, num_blocks=12,
                  vocab_size=32_000, causal=True)
paradigm = ARParadigm(cfg)
# hand to Trainer — paradigm.build_model() is called there

Every concrete paradigm accepts its architecture config positionally.

Declared here (not abstract) so callers can do type(some_paradigm)(new_config) — e.g. the W&B sweep agent rebuilding a paradigm with overridden hyperparameters — without a static type error, even though this base implementation is never actually used (each subclass overrides it; some, like EmbedderParadigm, add further keyword-only arguments with defaults, so calling with just config remains valid everywhere).

requires_shifted_targets: bool = True
build_model(rngs: Rngs) → Transformer[source]

Construct and return the NNX model for this paradigm.

Called once by the Trainer at the start of fit(). The returned model is then managed (checkpointed, sharded) by the Trainer and passed back into loss_fn / generate as the first argument.

loss_fn(model: Transformer, batch: Array, rng: Array, embeddings: Array | None = None) → tuple[Array, dict[str, Any]][source]

Compute the scalar training loss for one batch.

The model is passed explicitly so nnx.value_and_grad can differentiate through it without the paradigm needing to be an NNX module itself.

Parameters:

embeddings – Per-batch extras from prepare_batch, only passed when provides_batch_extras is True (e.g. pre-computed T5 embeddings for the continuous flow-matching paradigm). Ignored by paradigms that don’t declare provides_batch_extras = True.

Returns:

(scalar_loss, metrics_dict) where metrics_dict holds any auxiliary scalars (ce_loss, aux_loss, mse_loss, …) for logging.

generate(model: Transformer, prompt: Array, rng: Array, max_new_tokens: int = 200, temperature: float = 1.0, top_k: int | None = None, top_p: float | None = None, greedy: bool = False, use_cache: bool = True) → Array[source]

Generate token sequences given a prompt prefix.


Discrete Diffusion

class dantinox.paradigms.diffusion.discrete.DiscreteParadigm(model_config: ModelConfig, diffusion_config: DiscreteConfig | None = None)[source]

Bases: ParadigmBase

LLaDA-style masked-token diffusion paradigm.

Training objective: (1/t)-weighted cross-entropy on masked positions. Corruption: randomly mask tokens with probability p_mask(t) where t ~ Uniform(0, 1) per sample.

Quick-start:

cfg      = ModelConfig(dim=512, n_heads=8, num_blocks=12,
                       causal=False, noise_schedule="cosine")
# mask_token_id auto-detected from the tokenizer at Trainer.fit() time
paradigm = DiscreteParadigm(cfg)

A DiscreteConfig is still accepted as a second positional argument for backward compatibility.

Every concrete paradigm accepts its architecture config positionally.

Declared here (not abstract) so callers can do type(some_paradigm)(new_config) — e.g. the W&B sweep agent rebuilding a paradigm with overridden hyperparameters — without a static type error, even though this base implementation is never actually used (each subclass overrides it; some, like EmbedderParadigm, add further keyword-only arguments with defaults, so calling with just config remains valid everywhere).

property diffusion_config: Any

Return a lightweight namespace with noise_schedule and mask_token_id.

build_model(rngs: Rngs) → Transformer[source]

Construct and return the NNX model for this paradigm.

Called once by the Trainer at the start of fit(). The returned model is then managed (checkpointed, sharded) by the Trainer and passed back into loss_fn / generate as the first argument.

loss_fn(model: Transformer, batch: Array, rng: Array, embeddings: Array | None = None) → tuple[Array, dict[str, Any]][source]

Compute the scalar training loss for one batch.

The model is passed explicitly so nnx.value_and_grad can differentiate through it without the paradigm needing to be an NNX module itself.

Parameters:

embeddings – Per-batch extras from prepare_batch, only passed when provides_batch_extras is True (e.g. pre-computed T5 embeddings for the continuous flow-matching paradigm). Ignored by paradigms that don’t declare provides_batch_extras = True.

Returns:

(scalar_loss, metrics_dict) where metrics_dict holds any auxiliary scalars (ce_loss, aux_loss, mse_loss, …) for logging.

generate(model: Transformer, prompt: Array, rng: Array, max_new_tokens: int = 256, n_steps: int = 50, temperature: float = 1.0, decoding_strategy: str = 'sample', confidence_threshold: float = 0.9, factor: float = 1.5, block_size: int | None = None, steps_per_block: int = 50, use_dual_cache: bool = True, refresh_interval: int | None = None, verbose: bool = False) → Array[source]

Run reverse diffusion and return the generated token IDs.

Parameters:
  • decoding_strategy – "sample" (default), "greedy", "confidence", or "factor". When block_size is set only "threshold"/"confidence" and "factor" are meaningful (confidence-based strategies).

  • confidence_threshold – τ for the "confidence"/"threshold" strategy.

  • factor – f for the "factor" strategy.

  • block_size – If set, use block-wise generation (Fast-dLLM style). None (default) uses global denoising.

  • steps_per_block – Inner denoising steps per block (block mode only).

  • use_dual_cache – Enable prefix+suffix DualCache (block mode only).

  • refresh_interval – Recompute suffix cache every N steps (block mode only).

  • verbose – Print a per-step unmasking trace — tokens revealed, masks left, average confidence (global mode only).

stream(model: Transformer, prompt: Array, rng: Array, max_new_tokens: int = 256, n_steps: int = 50, temperature: float = 1.0, decoding_strategy: str = 'sample', confidence_threshold: float = 0.9, factor: float = 1.5, block_size: int | None = None, steps_per_block: int = 50, use_dual_cache: bool = True, refresh_interval: int | None = None, verbose: bool = False) → Iterator[tuple[int, int, Array]][source]

Yields (step, total, x_gen) after each denoising step.

Parameters:
  • decoding_strategy – "sample" (default), "greedy", "confidence", or "factor". When block_size is set only "threshold"/"confidence" and "factor" are meaningful.

  • confidence_threshold – τ for the "confidence"/"threshold" strategy.

  • factor – f for the "factor" strategy.

  • block_size – If set, use block-wise generation (Fast-dLLM style). None (default) uses global denoising.

  • steps_per_block – Inner denoising steps per block (block mode only).

  • use_dual_cache – Enable prefix+suffix DualCache (block mode only).

  • refresh_interval – Recompute suffix cache every N steps (block mode only).

  • verbose – Print a per-step unmasking trace — tokens revealed, masks left, average confidence (global mode only).


Continuous Flow-Matching

class dantinox.paradigms.diffusion.continuous.ContinuousParadigm(config: ModelConfig | FlowMatchingConfig)[source]

Bases: ParadigmBase

ELF (Embedded Language Flows) continuous flow-matching paradigm.

The forward process is z_t = t·x + (1−t)·ε where t ∈ [0,1], ε ~ N(0,I), and the model predicts the clean embedding x (x-prediction).

Architecture: a bidirectional transformer operating in a continuous embedding space, conditioned on in-context control tokens for timestep, CFG scale, and operating mode (denoiser vs. decoder branch).

Training requires a frozen T5 contextual encoder (transformers package, pip install dantinox[elf]). The encoder runs outside JIT; the Trainer obtains per-batch embeddings through prepare_batch and initialises the embedding normalisation statistics via on_train_start.

Quick-start:

cfg      = dx.ModelConfig(dim=512, n_heads=8, num_blocks=12,
                          embed_dim=768, bottleneck_dim=128, causal=False)
paradigm = ContinuousParadigm(cfg)

A raw FlowMatchingConfig is also accepted for Level-3 control over training hyper-parameters (denoiser schedules, CFG bounds, etc.).

Every concrete paradigm accepts its architecture config positionally.

Declared here (not abstract) so callers can do type(some_paradigm)(new_config) — e.g. the W&B sweep agent rebuilding a paradigm with overridden hyperparameters — without a static type error, even though this base implementation is never actually used (each subclass overrides it; some, like EmbedderParadigm, add further keyword-only arguments with defaults, so calling with just config remains valid everywhere).

provides_batch_extras: bool = True
build_model(rngs: Rngs) → FlowMatchingTransformer[source]

Construct and return the NNX model for this paradigm.

Called once by the Trainer at the start of fit(). The returned model is then managed (checkpointed, sharded) by the Trainer and passed back into loss_fn / generate as the first argument.

build_embedder(rngs: Rngs) → FlowEmbedder[source]

Build the frozen T5 embedder used to project tokens to flow space.

loss_fn(model: FlowMatchingTransformer, batch: Array, rng: Array, embeddings: Array | None = None) → tuple[Array, dict[str, Any]][source]

Compute the flow-matching training loss.

Parameters:
  • model – FlowMatchingTransformer NNX module.

  • batch – Integer token IDs [B, T] (targets for CE branch).

  • rng – JAX random key.

  • embeddings – Raw T5 contextual embeddings [B, T, embed_dim] from prepare_batch; normalised here via model.encode before the flow-matching loss.

Returns:

(scalar_loss, metrics_dict)

on_train_start(model: FlowMatchingTransformer, sample_batches: list[Any]) → None[source]

Initialise the embedder’s normalisation stats from real T5 outputs.

prepare_batch(batch: Any) → Array[source]

Run the frozen T5 encoder (outside JIT) → embeddings [B, T, E].

generate(model: FlowMatchingTransformer, prompt: Array | None = None, rng: Array | None = None, max_new_tokens: int | None = None, n_steps: int | None = None, cfg_scale: float | None = None, gamma: float | None = None, batch_size: int | None = None, seed: int | None = None) → Array[source]

Flow-matching generates unconditionally from Gaussian noise.

prompt only provides the batch size / sequence length defaults (max_new_tokens overrides its length); its token contents are unused. batch_size and seed can be passed directly as an alternative to providing a prompt and rng.

stream(model: FlowMatchingTransformer, prompt: Array | None = None, rng: Array | None = None, max_new_tokens: int | None = None, n_steps: int | None = None, cfg_scale: float | None = None, gamma: float | None = None) → Iterator[tuple[int, int, Array]][source]

Like generate but yields (step, total, tokens) after each ODE step.

num_parameters(model: FlowMatchingTransformer) → int[source]

Count trainable parameters in the model.


Sentence Embedder

class dantinox.paradigms.embedder.EmbedderParadigm(config: ModelConfig, *, pooling: str = 'auto', temperature: float = 0.05)[source]

Bases: ParadigmBase

Contrastive embedding paradigm (SimCSE unsupervised).

Works directly with Trainer — no labelled pairs required. Each [B, T] token batch is encoded twice with different dropout masks; the two views are treated as anchor/positive and trained with InfoNCE.

Quick-start:

import dantinox as dx

cfg = dx.ModelConfig(
    dim=256, n_heads=4, head_size=64, num_blocks=4,
    vocab_size=32_000, causal=False,   # bidirectional → best embeddings
)
paradigm = dx.EmbedderParadigm(cfg)

# train with the stock Trainer — any flat text corpus works
run_dir = dx.train(paradigm, "data/corpus.txt", lr=3e-4, epochs=5)

# use as encoder
embedder = dx.Embedder.from_run(run_dir)
vecs = embedder.embed(["hello world", "foo bar"])

Parameters

config: Model architecture config (prefer causal=False for better

bidirectional representations).

pooling: "auto" | "mean" | "last" | "cls". temperature: InfoNCE temperature (default 0.05).

Every concrete paradigm accepts its architecture config positionally.

Declared here (not abstract) so callers can do type(some_paradigm)(new_config) — e.g. the W&B sweep agent rebuilding a paradigm with overridden hyperparameters — without a static type error, even though this base implementation is never actually used (each subclass overrides it; some, like EmbedderParadigm, add further keyword-only arguments with defaults, so calling with just config remains valid everywhere).

requires_shifted_targets: bool = False
build_model(rngs: Rngs) → Transformer[source]

Construct and return the NNX model for this paradigm.

Called once by the Trainer at the start of fit(). The returned model is then managed (checkpointed, sharded) by the Trainer and passed back into loss_fn / generate as the first argument.

loss_fn(model: Transformer, batch: Array, rng: Array, embeddings: Array | None = None) → tuple[Array, dict[str, Any]][source]

SimCSE: encode the same batch twice with dropout → InfoNCE.

generate(model: Transformer, prompt: Array, rng: Array, **kwargs: Any) → Array[source]

For embedders generate returns pooled embeddings of prompt.


Usage examples

Standard AR training

import dantinox as dx
from flax import nnx

cfg      = dx.ModelConfig(paradigm="ar", dim=256, n_heads=8, head_size=32,
                           num_blocks=6, vocab_size=200)
paradigm = dx.Paradigm(cfg)
model    = paradigm.build_model(nnx.Rngs(42))

print(f"Parameters: {paradigm.num_parameters(model) / 1e6:.2f}M")
print(f"Type: {paradigm.type}")   # → "ar"

# Single train step
import jax.numpy as jnp
batch = {"input_ids": jnp.ones((4, 64), dtype=jnp.int32)}
loss, metrics = paradigm.loss_fn(model, batch, rng=nnx.Rngs(0))

Discrete diffusion

import dantinox as dx

cfg      = dx.ModelConfig(
    paradigm="discrete",          # causal=False auto-configured
    dim=256, n_heads=8, head_size=32,
    num_blocks=6, vocab_size=32000,
    noise_schedule="cosine",
    mask_token_id=4,
)
paradigm = dx.Paradigm(cfg)
model    = paradigm.build_model(nnx.Rngs(42))

Continuous flow-matching

import dantinox as dx

cfg = dx.ModelConfig(
    paradigm="continuous",        # causal=False auto-configured
    embed_dim=768,                # must match T5-base hidden size
    bottleneck_dim=128,
    dim=512, n_heads=8, head_size=64,
    num_blocks=6, vocab_size=32128,
)
paradigm = dx.Paradigm(cfg)
model    = paradigm.build_model(nnx.Rngs(42))
embedder = paradigm.build_embedder()   # frozen T5 oracle

Sentence embedder (InfoNCE)

import dantinox as dx

cfg = dx.ModelConfig(
    paradigm="embedder",          # causal=True auto-configured
    dim=512, n_heads=8, head_size=64,
    num_blocks=12,
    embed_pooling="mean",         # "mean" | "last" | "cls" | "auto"
    embed_temperature=0.05,
)
paradigm = dx.Paradigm(cfg)

See also