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_mxto 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.pyaccept--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 --validateenables 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¶
config.py, decode_metadata.py, and rope_tables.py have no __main__
block and are imported rather than run.