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: true

This 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_proj

in_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.