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:
Model construction —
build_model()returns the JAX/NNX model for this paradigm.Loss function —
loss_fn(model, batch)returns(loss, metrics). TheTrainernever touches the model directly — it calls this.Generation —
generate()wraps the model-specific decode loop.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
|
Implementation |
|
Model class |
|---|---|---|---|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
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:
ParadigmBaseUnified 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
paradigminModelConfigalso auto-configurescausal, so you never need to writecausal=Falsefor 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, likeEmbedderParadigm, add further keyword-only arguments with defaults, so calling with justconfigremains 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 intoloss_fn/generateas 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_gradcan differentiate through it without the paradigm needing to be an NNX module itself.- Parameters:
embeddings – Per-batch extras from
prepare_batch, only passed whenprovides_batch_extrasis True (e.g. pre-computed T5 embeddings for the continuous flow-matching paradigm). Ignored by paradigms that don’t declareprovides_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_extrasis True; the return value is forwarded toloss_fnas theembeddingskeyword argument.
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:
ABCAbstract 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_fnandgenerate— 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 toloss_fnas theembeddingskeyword. Setprovides_batch_extras = Trueto 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, likeEmbedderParadigm, add further keyword-only arguments with defaults, so calling with justconfigremains valid everywhere).- __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, likeEmbedderParadigm, add further keyword-only arguments with defaults, so calling with justconfigremains 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 intoloss_fn/generateas 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_gradcan differentiate through it without the paradigm needing to be an NNX module itself.- Parameters:
embeddings – Per-batch extras from
prepare_batch, only passed whenprovides_batch_extrasis True (e.g. pre-computed T5 embeddings for the continuous flow-matching paradigm). Ignored by paradigms that don’t declareprovides_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.
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:
ParadigmBaseAutoregressive 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, likeEmbedderParadigm, add further keyword-only arguments with defaults, so calling with justconfigremains valid everywhere).- 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 intoloss_fn/generateas 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_gradcan differentiate through it without the paradigm needing to be an NNX module itself.- Parameters:
embeddings – Per-batch extras from
prepare_batch, only passed whenprovides_batch_extrasis True (e.g. pre-computed T5 embeddings for the continuous flow-matching paradigm). Ignored by paradigms that don’t declareprovides_batch_extras = True.- Returns:
(scalar_loss, metrics_dict) where metrics_dict holds any auxiliary scalars (ce_loss, aux_loss, mse_loss, …) for logging.
Discrete Diffusion
- class dantinox.paradigms.diffusion.discrete.DiscreteParadigm(model_config: ModelConfig, diffusion_config: DiscreteConfig | None = None)[source]
Bases:
ParadigmBaseLLaDA-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
DiscreteConfigis 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, likeEmbedderParadigm, add further keyword-only arguments with defaults, so calling with justconfigremains 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 intoloss_fn/generateas 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_gradcan differentiate through it without the paradigm needing to be an NNX module itself.- Parameters:
embeddings – Per-batch extras from
prepare_batch, only passed whenprovides_batch_extrasis True (e.g. pre-computed T5 embeddings for the continuous flow-matching paradigm). Ignored by paradigms that don’t declareprovides_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". Whenblock_sizeis 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". Whenblock_sizeis 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:
ParadigmBaseELF (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 (
transformerspackage,pip install dantinox[elf]). The encoder runs outside JIT; the Trainer obtains per-batch embeddings throughprepare_batchand initialises the embedding normalisation statistics viaon_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
FlowMatchingConfigis 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, likeEmbedderParadigm, add further keyword-only arguments with defaults, so calling with justconfigremains valid everywhere).- 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 intoloss_fn/generateas 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]fromprepare_batch; normalised here viamodel.encodebefore 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_tokensoverrides its length); its token contents are unused.batch_sizeandseedcan 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
generatebut yields(step, total, tokens)after each ODE step.
Sentence Embedder
- class dantinox.paradigms.embedder.EmbedderParadigm(config: ModelConfig, *, pooling: str = 'auto', temperature: float = 0.05)[source]
Bases:
ParadigmBaseContrastive 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=Falsefor 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, likeEmbedderParadigm, add further keyword-only arguments with defaults, so calling with justconfigremains valid everywhere).- 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 intoloss_fn/generateas the first argument.
- config: Model architecture config (prefer
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
Configuration Reference — all config fields including
paradigm,embed_pooling,embed_temperatureArchitecture: Paradigm System — conceptual explanation
Training API — how
Trainercallsloss_fn