Skip to content

DeepSeek V4-Pro

models/deepseek_v4_pro/ is the Ascend 950 (A5) implementation for DeepSeek-V4 Pro, with an optional Flash architecture preset. Select the preset at import and compile time with DEEPSEEK_V4_VARIANT=pro|flash or --variant pro|flash.

Deployment configuration

The PRO and FLASH presets in config.py define the architecture-specific shapes and layer schedules. Pro remains the default so existing operator entry points and DailyCI keep their prior behavior.

Deployment property Value
Speculative decoding MTP = 1 (DECODE_SEQ = 2)
Decode batch per card 4 requests → 8 token rows per step
Decode context length up to 16,384 positions (KERNEL_MAX_SEQ_LEN), 128-token pages
Prefill shape one request of 128 tokens per program
Platform Ascend A5 (-p a5); full forwards are device-only
Expert parallelism --ep 2/4/8; real Flash requires EP8 with 32 routed experts per rank
LM-head parallelism --tp 2/4/8 vocab shards; LM_HEAD_TP_SIZE = 8 is the deployment value
Other components no tensor parallelism — attention is data-parallel, MoE is expert-parallel
Quantization Hybrid MXFP8-MXFP4 — MXFP8 for the dense path, MXFP4 for the routed-expert weights
Serving none — no pypto-serving deployment consumes this tree

PRO.max_position_embeddings keeps the architectural one-million-position value, but admitting 1 M positions would need a ~64× larger physical pool than the cases allocate and a host-side golden nobody can compute. PRO_KERNEL therefore replaces it with KERNEL_MAX_SEQ_LEN = 16384 — an 8k prompt plus 512 decode steps, the budget the Flash cases already exercise. Raise that one constant if a case needs a longer context.

The native MX path follows the CANN DeepSeek V4 deployment quantization boundary:

  • Routed W1/W3/W2 tensors remain packed MXFP4 on disk and in device memory. utils.py canonicalizes source groups to the CANN MXFP4 weight representation and packs their per-32 E8M0 scales as MX_B_NN. expert_routed.py expands each E2M1 nibble exactly to FP8E4M3 per tile before Cube multiplication; this device bridge is lossless.
  • Shared W1/W3/W2 use native MXFP8 data and per-32 E8M0 scales.
  • gate.py applies pl.quant_mx to the normalized MoE input. moe.py dispatches both the FP8 data and its scale. Both expert kernels round W1/W3 outputs to BF16, apply clipped SwiGLU, then use the CANN SwiGLU-specific ceiling exponent before W2. This scale rule differs from ordinary dynamic MX activation quantization.

The shared-expert kernels pass ordinary pl.load results directly to pl.matmul_mx. The routed experts do not: their W1/W3/W2 payloads are packed MXFP4 in device memory and are expanded to MXFP8 through a vector LUT gather before the Cube op. Either way PyPTO infers the data and scale staging from the four operand positions, including LeftScale and RightScale placement. Quantized activations remain GM-backed where multiple expert or W2 output blocks reuse them; direct vector-to-cube transport would otherwise repeat quantization or reduce output-block parallelism.

This path does not use an NVIDIA scale_alg setting. Scale generation and physical layout conversion are expressed directly through PyPTO's MX APIs.

Model shape and layer schedule

Pro is wider and deeper than Flash: 7168 hidden, 128 attention heads, a 1536 Q-LoRA rank, 16 output-projection groups, moe_intermediate_size = 3072, 384 routed experts with top-6 routing plus 1 shared expert, and an indexer top-k of 1024. compress_ratios carries 62 entries — 61 model layers plus the MTP layer:

Ratio Path Layers Count
128 HCA — ratio-128 compressor, deterministic top-k 0, 1, 3, 5, …, 59 31
4 CSA — ratio-4 compressor plus the learned indexer 2, 4, …, 60 30
0 SWA — sliding window only the MTP layer 1

Unlike Flash, the main stack has no SWA layer; SWA appears only in the MTP layer, which is why decode_attention_swa is reachable from decode_mtp but not from decode_fwd. The first three layers route by hash (num_hash_layers = 3).

Model structure, top down

decode_fwd and prefill_fwd

Both hand-unroll the layer schedule inside one rank-generic @pl.jit kernel launched per EP rank, with every attention and MoE stage in its own pl.scope():

decode_fwd     layers 0, 1        decode_attention_hca → moe
               loop over pairs    decode_attention_csa → moe   (even layers)
                                  decode_attention_hca → moe   (odd layers)
               layer 60           decode_attention_csa → moe
               tail               hc_head → rms_norm

prefill_fwd    same schedule with prefill_attention_{hca,csa} → moe,
               tail               hc_head → rms_norm

Both forwards finish with the final norm and LM-head sampling. The standalone lm_head.py entry point validates that distributed tail separately.

One layer

attention   hc_pre → rmsnorm → qkv_proj_rope → (compress / index) → sparse_attn → hc_post
moe         hc_pre → gate → expert_shared → dispatch → expert_routed → combine → hc_post

decode_layer and prefill_layer expose exactly this pair as standalone two-rank harnesses.

Attention paths

decode_attention_swa   hc_pre → rmsnorm → qkv_proj_rope
                              → decode_sparse_attn_swa                   → hc_post
decode_attention_hca   hc_pre → rmsnorm → qkv_proj_rope
                              → decode_compressor_ratio128
                              → decode_sparse_attn_hca                   → hc_post
decode_attention_csa   hc_pre → rmsnorm → qkv_proj_rope
                              → decode_compressor_ratio4 (main + inner)
                              → decode_indexer → decode_indexer_compressor
                              → decode_sparse_attn                       → hc_post

prefill_attention_swa  hc_pre/hc_post + rmsnorm + qkv_proj_rope + prefill_sparse_attn
prefill_attention_hca  … + prefill_compressor_ratio128
prefill_attention_csa  … + prefill_compressor_ratio4
                          + prefill_indexer → prefill_indexer_compressor

decode_sparse_attn is the CSA variant here (the Flash tree names it decode_sparse_attn_csa); all three own the fused grouped output projection. rope_tables generates the RoPE/YaRN tables on the host and decode_metadata lowers the fixture's paged-cache metadata.

MTP path

mtp_projection  e_proj(enorm(hidden)) + h_proj(hnorm(prev_hidden))
decode_mtp      mtp_projection → decode_attention_swa → moe → hc_head → rmsnorm
prefill_mtp     mtp_projection → prefill_attention_swa → moe → hc_head → rmsnorm

prefill_mtp reuses prefill_fwd's driver for the main-model pass.

Real weights (Flash)

utils.py converts the released DeepSeek-V4-Flash checkpoint into the native MX tensor ABI of the two forward drivers. Dense attention projections, the indexer query projection, and shared experts use MXFP8 with per-32 E8M0 scales. Routed experts retain packed MXFP4 payloads with per-32 E8M0 scales. Per-layer tensors are stacked and EP/TP-sharded exactly like the fixture specs.

Reference and numerical boundaries

The architecture and checkpoint source are the official DeepSeek Flash release. The A5 quantization reference is CANN recipes. The cache records both revisions. The implementation targets this CANN deployment convention; it does not promise bitwise identity with the released Triton inference path, whose activation grouping and indexer quantization differ.

Boundary A5 implementation
Dense/shared weights Source block-FP8 dequantization, then CANN MXFP8 KN groups of 32
Routed weights CANN MXFP4 group conversion; exact E2M1-to-E4M3 device expansion
Ordinary linear input BF16 input, OCP dynamic MXFP8 groups of 32
Expert intermediate BF16 W1/W3 outputs, clipped SwiGLU, ceiling MXFP8 scale before W2
Routed reduction BF16 W2 output, then routing multiplier
Q projection BF16 linear/norm boundaries; Q-head RMS norm has no learned affine weight
Main KV Non-RoPE channels undergo FP8 quantize/dequantize in groups of 64; RoPE channels bypass quantization
Indexer Q BF16 Hadamard boundary, E4M3 with a linear per-head scale
Indexer K BF16 Hadamard boundary, E4M3 with a power-of-two per-row ceiling scale
Residual BF16 rounding at the hyperconnection post boundary

Flash uses 43 layers, 4096 hidden channels, 64 attention heads, 256 routed experts, top-6 routing, and one shared expert. Layers 0 and 1 use SWA; the remaining main layers alternate ratio-4 CSA and ratio-128 HCA, ending with CSA. The first three layers use hash routing. The local context limit is 16,384 positions, below the checkpoint's architectural one-million limit.

A completed cache has a schema-versioned manifest.json containing the source, quantization reference, EP/TP configuration, and each tensor's shape, dtype, byte size, and SHA-256. Readers reject stale, incomplete, or corrupt caches. Conversion writes tensors atomically and publishes the manifest last.

Convert once offline, then point the drivers at the cache:

PYTHONPATH=.:models/deepseek_v4_pro python -c 'import utils; utils.main()' \
    --variant flash --ep 8 --tp 2 \
    --ckpt /path/to/DeepSeek-V4-Flash --out build_output/flash_weights_ep8_tp2
python models/deepseek_v4_pro/prefill_fwd.py --variant flash --ep 8 --tp 2 \
    -p a5 -d 0,1,2,3,4,5,6,7 --weights build_output/flash_weights_ep8_tp2

--weights also accepts the raw checkpoint directory (converted on the fly; slower and RAM-hungry — the cache is the recommended path). Only EP8 deploys the full 256-expert model: the kernel programs keep 32 local experts per rank (moe.py shrinks the global routing space to 32*EP), so an EP4/EP2 real-weight run uses the first 32*EP checkpoint experts with reduced router tables — a smoke configuration, not the true model output.

Numeric validation on real weights:

  • decode_layer.py / prefill_layer.py accept --weights <ckpt_dir> to inject one layer's real weights (converted on demand); the layer golden then recomputes with the same weights, so the existing per-layer validation runs on real dynamic ranges.
  • prefill_fwd.py --validate enables a full-network torch golden (utils.golden_prefill_fwd: embed → 43 chained layer goldens → hc_head → final norm → LM head). End-of-network gates are cosine/rel-L2 on the selected logit rows plus greedy-sample agreement; per-element gates on deep hidden states and compressor state pools accumulate cross-layer drift and are expected to need looser budgets than the single-layer drivers. RoPE tables, the indexer Hadamard, caches, and per-step metadata keep their fixture initializers. Without --validate, a real-weight forward is a runtime smoke test; numerical agreement requires the explicit golden comparison.

Golden data can be computed once and replayed: prefill_fwd.py --validate --save-data persists the generated inputs and golden outputs under <runtime_dir>/data/, and --golden-data <dir> loads them back and runs only the device pass plus the comparison — the CPU-heavy golden compute can run on a host without NPU access while the short device pass reuses it. --prompt-file <file> --tokenizer <tokenizer.json> replaces the synthetic input_ids with a real prompt (replicated across ranks; num_tokens follows the prompt length).

End-to-end token generation

synthetic_token_loop.py drives the full prompt-to-text path on real weights: the prompt is encoded with the checkpoint's tokenizer.json (BOS prepended unless --no-bos), the resident session runs one prefill plus --decode-steps greedy decode steps, every step asserts that all ranks sampled the same token, and the sampled ids are detokenized at the end. Decoding stops early when --eos-id (default 1) is sampled. Without --weights the loop keeps its synthetic zero-weight control-path behavior.

The EP8 example uses the fixed 128-row prefill capacity; active rows follow the encoded prompt length:

python models/deepseek_v4_pro/synthetic_token_loop.py --variant flash \
    --ep 8 --tp 2 -d 0,1,2,3,4,5,6,7 \
    --weights build_output/flash_weights_ep8_tp2 \
    --tokenizer /path/to/DeepSeek-V4-Flash/tokenizer.json \
    --prompt "The capital of France is" --decode-steps 32 \
    --result-json build_output/flash_e2e.json

The EP8 resident session prepares about 210 GiB of host tensors in shared memory. Ensure the Linux /dev/shm mount has enough free space before starting; the loader checks capacity before promoting the resident bank. This host preparation is included in load time and excluded from token-generation timing.

The optional result JSON records success/failure, generated token IDs/text, compile and load time separately, resident time to first token, and inter-token wall times. Throughput counts the generated sequence once; EP ranks replicate the request. Timings include host coordination and are measured without warmup, so they are not serving throughput or isolated kernel latency.

It also records rank_spread, one entry per sampling point, and the ranks_bit_identical verdict over the whole run. Every rank holds the full hidden state and the full vocabulary after the MoE combine and the LM-head gather, so the ranks should be bit-identical; equal sampled tokens do not show that, because drift only becomes a token difference when it flips a near-tie argmax. A run can therefore report pass with ranks_bit_identical false. The spread is scanned after each step is timed, so it is outside every latency figure.

The full prefill and decode programs carry runtime num_tokens and moe_epoch_base scalars in their compiled ABI. Their ScalarSpecs use compile_runtime=True, so run passes pl.RUNTIME during signature-driven compilation instead of folding the initial values into generated task arguments. num_tokens follows the real prompt/decode row count, while callers advance the epoch scalar by LAST_MOE_EPOCH for every physical dispatch on a persistent worker.

MoE payload readiness uses one cache-line-padded epoch slot per source and producer block. Each dispatch or combine block stores its current epoch with Set only after that block's self-draining tensor puts; a separate whole-grid wait observes every remote slot with >= epoch before gather or reduction. A separate per-rank consumed epoch is published after the complete reduction, and the next MoE invocation waits for every rank to consume the previous epoch before reusing payload windows. This avoids shared-counter atomic fan-in, detached notifications, and mixed readiness/lifetime credit arithmetic. The full forwards, standalone MoE, and decode-layer driver keep these epochs monotonic for the lifetime of their persistent program. The packed prefill-layer and fixed-epoch MTP drivers instead quiesce all final consumed markers and clear only their own inbound slots before a later synchronous dispatch.

Both full programs must be recompiled when moving from an older artifact. The token loop rejects artifacts that omit the runtime scalar or whose generated host_orch.py does not forward it through TaskArgs.add_scalar. After a timeout or partial dispatch failure, callers must discard the worker and its persistent windows rather than retrying with a guessed epoch.

The prefill and decode RoPE paths use fixed even/odd lane gather and scatter operations for adjacent-lane permutations instead of synthesizing tile-local index tensors.

The EP8 loop needs a toolchain without the pto-isa A5 dispatch regression that stalls every EP8 prefill before the first token, and a probabilistic cross-rank divergence (#1043) is still open, so a run can diverge between ranks without the kernels having changed. Read ranks_bit_identical rather than the pass/fail status when judging whether a given run was affected.

What the nightly end-to-end job checks

The e2e-flash-a5 job in daily_ci.yml runs this loop on the pinned checkpoint and publishes the prompt, the generated text and the timings in the run summary. It asserts only that generation happened: eight distinct devices, a token count that matches the completed decode steps, and decoded text with lexical content. Whether the text is good is for a reader of the summary to judge, so an operator change that damages the model shows up there as degraded output, or below as a failure. It has no continue-on-error and no skip path: a runner without PYPTO_DSV4_FLASH_CKPT_DIR fails the job rather than reporting a green night that never sampled a token.

The converted weight cache is reused across nights only while _load_complete_cache still accepts it against the checked-out conversion code, so a change to the quantization path is converted again instead of being masked by the previous night's tensors.

Files

Group Files
Full forward decode_fwd.py, prefill_fwd.py
Layer composition decode_layer.py, prefill_layer.py
MTP decode_mtp.py, prefill_mtp.py, mtp_projection.py
Decode attention orchestration decode_attention_swa.py, decode_attention_csa.py, decode_attention_hca.py
Decode sparse attention (fused o-proj) decode_sparse_attn.py, decode_sparse_attn_swa.py, decode_sparse_attn_hca.py
Decode compressors and indexer decode_compressor_ratio4.py, decode_compressor_ratio128.py, decode_indexer.py, decode_indexer_compressor.py
Prefill attention and cache prefill_attention_swa.py, prefill_attention_csa.py, prefill_attention_hca.py, prefill_sparse_attn.py, prefill_compressor_ratio4.py, prefill_compressor_ratio128.py, prefill_indexer.py, prefill_indexer_compressor.py
Shared transforms rmsnorm.py, qkv_proj_rope.py, hc_pre.py, hc_post.py, hc_head.py
MoE and output moe.py, gate.py, expert_shared.py, expert_routed.py, lm_head.py
Metadata and host helpers config.py, decode_metadata.py, rope_tables.py, utils.py
Real-weight loading and MX conversion utils.py
Token loop synthetic_token_loop.py

config.py, decode_metadata.py, and rope_tables.py have no __main__ block and are imported rather than run.