AREX

AREX-2 27B · INT8 W8A16 · MTP

Target-adapted multi-token prediction for speculative decoding.

Base model · Paper · Project · vLLM

Format W8A16 Weights INT8 MTP depth 3 Activations BF16 License Apache 2.0

Model description

This model combines the BAAI/AREX-2 27B text-generation checkpoint with an adapted multi-token prediction (MTP) module for speculative decoding. The AREX-2 target uses compressed-tensors symmetric group-128 W8A16 weight-only quantization; its embeddings and output head are retained in BF16, as are activations. The added MTP module and its draft head are also BF16.

During speculative decoding, the MTP module proposes up to three tokens and the AREX-2 target verifies them. The target weights were kept frozen during MTP adaptation; the MTP module is an auxiliary draft component, not a replacement language model.

Architecture and provenance

  • Target model: BAAI/AREX-2, source revision d4e3502f92d9e889031c2d04e387ab2eb3520268.
  • MTP initialization: Qwen3.8 BF16 checkpoint revision 1d4bf0f2ff6012fd82039f2fa52739d0dd7c60c0; its 15 BF16 MTP tensors were used only to initialize the draft module.
  • Initial MTP tensor SHA-256: 1d8268aa85ace093a561e3e7b63b9d390dac1cd55a90cd55b5ec509c3c9da9fe.
  • Trained MTP tensor SHA-256: 38892896461c1a2c30609aebfba9b99ab0eb57a8cb49830ab98b817a398e924a.
  • The target and initializer share the same text configuration: Qwen3_5, hidden size 5,120, 64 layers, vocabulary size 248,320. The AREX-2 tokenizer is used because its pre-tokenizer and decoder settings differ from Qwen3.8.
  • The MTP draft head uses a 40,960-token shortlist, selected primarily from tokens observed in AREX-2-generated continuations. Held-out shortlist coverage was 97.5732% of generated target tokens.

Intended use and limitations

Intended for text generation, code assistance, and reasoning with speculative decoding through a compatible vLLM MTP implementation. The AREX-2 model is the target; the added MTP component proposes draft tokens and is not a standalone language model.

The W8A16 target is quantized and may differ numerically from its BF16 source. The distillation set is predominantly English code and math reasoning, with a smaller mixed-language supplement. The offline held-out metrics below do not establish general task quality or a serving-speed improvement for a specific workload.

Distillation data and method

AREX-2 W8A16 was frozen and generated the continuations used to capture teacher hidden states and output distributions. Source-dataset reference answers were not used; no private user or agent-session data was used.

  • 1,000 English prompts: 600 code-instruction prompts from Magicoder OSS-Instruct-75K and 400 math-reasoning questions from the GSM8K train split.
  • 228 additional target-generated examples from a mixed prompt pool: 82 UltraChat, 39 Magicoder, 15 GSM8K, 36 Danish instruction, 28 Danish reasoning, and 28 skoleGPT examples.
  • Combined set: 1,228 sequences, 304,416 prompt tokens, and 468,614 AREX-2-generated continuation tokens; 773,030 hidden-state token rows.
  • The 40,960-token shortlist prioritizes observed AREX-2 output tokens and is filled from HyperQwen's compatible shortlist. Held-out output-token coverage in shortlist construction was 97.5732%.
  • Training unrolls MTP depth 3 with depth weights 1, 0.5, 0.25; the target output head is used for the target distribution, and the draft shortlist head is optimized jointly with the MTP module.
  • One epoch: 90 optimizer steps, 8,192 tokens per step, 512-token micro-batches, MTP learning rate 2e-5, draft-head learning rate 1e-5, and a 5-step warmup.
  • Training used 1,167 sequences and 738,390 target tokens; 61 sequences were held out for evaluation. The AREX-2 target weights remained frozen.
  • Training completed on a single 24-GiB RTX 3090 without out-of-memory errors.
  • Training implementation: HyperQwen MTP trainer, source commit e1459c7631774f56de2f9425437d54e7e72ea688.

Held-out evaluation

Before adaptation, the Qwen3.8-initialized MTP module achieved the following on held-out AREX-2 continuations, compared with the trained module. Top-1 is the draft token's agreement with the target's most likely token; acceptance is the simulated target-verified draft acceptance rate; KL is measured in nats per token.

Draft position Initial top-1 Trained top-1 Initial accepted Trained accepted Initial KL (nats/token) Trained KL (nats/token)
1 82.39% 83.02% 87.02% 87.37% 0.2989 0.2280
2 76.57% 77.86% 79.49% 79.66% 0.6742 0.4749
3 72.15% 74.10% 74.68% 74.80% 0.9813 0.6456

Held-out chain simulation increased from 2.837 to 2.884 accepted draft tokens per step (7,402 validation steps). These are offline self-distillation metrics, not a task-quality or end-to-end speed benchmark.

Inference

vLLM compatibility

This checkpoint uses a 40,960-token shortlist draft head, stored as mtp.draft_lm_head.weight, with its token IDs in mtp_draft_vocab_ids.pt. This is not the standard full-vocabulary Qwen3.5 MTP head layout. Stock vLLM builds that do not support this draft-vocabulary format can fail while loading the MTP weights because Qwen3_5MultiTokenPredictor has no draft_lm_head module.

Native MTP therefore requires a compatible vLLM build with draft-vocabulary support. The HyperQwen Qwen3.5 MTP draft-vocabulary patch is a reference implementation; check its compatibility with your vLLM revision before applying it. The serving image used for this model includes matching support. Stock vLLM 0.30.0 and the tested 0.30.1rc1 nightly builds have been reported to fail with this checkpoint. Changing the speculative method to qwen3_next_mtp does not address this weight-loading incompatibility; method: "mtp" is the correct method for this architecture.

With a compatible runtime, configure three speculative tokens:

vllm serve <model-id> \
  --trust-remote-code \
  --speculative-config '{"method":"mtp","num_speculative_tokens":3}'

Use the tokenizer shipped with this checkpoint. The reported MTP values are offline held-out measurements; serving throughput and task quality depend on the inference engine, batch/queue behavior, prefix-cache reuse, and workload.

Acknowledgements and license

This derivative is based on BAAI/AREX-2 and uses Qwen3.8 MTP weights only for initialization. Prompt sources include Magicoder OSS-Instruct-75K, GSM8K, and the supplemental instruction/reasoning datasets listed above. Model and dataset attribution and license terms remain with their respective authors; consult each source card before redistribution or commercial use. The base model is Apache-2.0 licensed.

Downloads last month
517
Safetensors
Model size
28B params
Tensor type
I32
·
BF16
·
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support

Model tree for numsu/AREX-2-27B-INT8-W8A16-MTP

Base model

Qwen/Qwen3.8-27B
Finetuned
BAAI/AREX-2
Quantized
(7)
this model

Paper for numsu/AREX-2-27B-INT8-W8A16-MTP