AREX-2 27B · INT8 W8A16 · MTP
Target-adapted multi-token prediction for speculative decoding.
Base model · Paper · Project · vLLM
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 rate1e-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