FLA Mamba and Mamba2
Mamba checkpoints use Transformers by default. To select the FLA implementation explicitly, set the backend in the model configuration:
base_model: state-spaces/mamba-130m-hf
model_config:
mamba_backend: fla
sample_packing: trueThis path supports full fine-tuning and LoRA on in_proj and out_proj for the
mamba and mamba2 model types. For LoRA, set:
adapter: lora
lora_target_modules:
- in_proj
- out_projin_proj adapters run before the fused kernels. When out_proj is adapted,
training uses FLA’s decomposed CUDA convolution/scan path so the output projection
invokes the adapter, including its dropout, rather than reading the base weight
directly. This can have different performance from full fine-tuning’s fused path.
Mamba2 with multiple state groups also uses the decomposed path to preserve
Transformers’ normalization across the full intermediate dimension. Previously
saved full-FLA Mamba2 checkpoints retain their original grouped normalization.
Other adapter targets and quantized model weights (including QLoRA and 8-bit
LoRA) are not supported. Validation rejects quantized loading and quantization
metadata in saved checkpoints. Use an unquantized checkpoint for this backend;
there is no automatic QLoRA fallback to Transformers because its fused paths
also access projection weights directly.
ModelSupport registers config-selected replacements for the pure Mamba mixers. Transformers retains the outer model, parameter initialization, checkpoint keys, loss, and generation interfaces. Hybrid models are not patched. Saved checkpoints retain the backend and remain loadable through Axolotl’s model loader.
Install the fla extra and use Torch 2.13 or newer. The wrapper loads the
kernels-community/mamba-ssm and kernels-community/causal-conv1d Hub kernels
when preparing CUDA computation. CUDA convolution is required; leave
FLA_CONV_BACKEND unset or set it to cuda.
Packing groups documents into batches with bounded right padding. Convolution
and recurrent state start independently for each document. Packing and Trainer
loss normalization belong to Axolotl core. Packed inputs cannot use a generation
cache; ordinary unsharded generation uses Transformers’ cache_params interface with
single-token cached decoding.
Packed training requires serial forward/backward calls on each model instance. Overlapping or interleaved calls can overwrite the document boundaries retained for gradient checkpointing.
Context parallelism requires Ringmaster 0.2.2 or newer for its FLA Mamba adapters. Older releases without those adapters are rejected. CP uses contiguous shards and exchanges convolution halos and recurrent states; packed document boundaries remain active across shard boundaries. Cached generation is unavailable while CP adapters are installed.
The 32K/64K BF16 microbenchmarks used four layers, hidden size 768, and batch size two. They exercised dense forward and backward, excluding compilation and optimizer work. Differences from Transformers remained bounded but were not bitwise zero; the results do not establish full training equivalence.