Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
52 commits
Select commit Hold shift + click to select a range
4b8b9a3
opt(deepseek-v4 projection): re-capture flash traces cleanly + recali…
lhzhang333 Jun 22, 2026
67b6b7e
ci(deepseek-v4 projection): deploy from dev/tas/deepseek-v4 branch
wenxie-amd Jun 22, 2026
6155507
DSv4: support in the performance-projection benchmark (#776)
JohnQinAMD Jun 23, 2026
0cfd6ba
DSv4: FP8/FP4 low-precision training recipe + Muon mxfp8 fix + launch…
JohnQinAMD Jun 23, 2026
e89b4c0
DSv4: FlyDSL attention backend + MI355X Triton tuning (#778)
JohnQinAMD Jun 23, 2026
aaa4f27
Dev/john/dsv4 gfx1250 port (#785)
JohnQinAMD Jun 25, 2026
159c196
feat: bump primus turbo wrapper to latest (#786)
RuibinCheung Jun 25, 2026
cb9510d
v4-pro launcher: default-ON the tested Triton small-kernel fusions
JohnQinAMD Jun 25, 2026
135ec6f
feat: add pad_offsets to match latest TE
RuibinCheung Jun 25, 2026
574b9aa
fix(deepseek-v4 projection): correct optimizer sharding model
lhzhang333 Jun 25, 2026
5b09425
opt(deepseek-v4 projection): model MTP flops and timing
lhzhang333 Jun 25, 2026
2390d56
opt(deepseek-v4 projection): support pipeline layout and selective re…
lhzhang333 Jun 25, 2026
cdba154
fix(deepseek-v4 projection): improve trace attribution and capture pa…
lhzhang333 Jun 25, 2026
671d417
opt(deepseek-v4 projection): modify trace location
lhzhang333 Jun 26, 2026
861600a
update dsv4 flash script
lhzhang333 Jun 26, 2026
a9e9d65
fix(deepseek-v4 projection): validate projection UI controls
lhzhang333 Jun 26, 2026
2c9b4c2
opt(deepseek-v4 projection): sticky model/gpu switcher bar
lhzhang333 Jun 26, 2026
032e87b
gfx1250 mxfp8: skip expert token zero-padding to avoid quant NaN
JohnQinAMD Jun 28, 2026
02b0871
muon: bump emerging_optimizers to 0.4.0a0 + grouped-expert orthogonal…
JohnQinAMD Jun 28, 2026
7137145
v4-pro muon launcher: gfx1250 enablement, gzip traces, compile-disabl…
JohnQinAMD Jun 28, 2026
901facb
feat(muon): batched grouped-expert Newton-Schulz orthogonalize (TP=1)…
JohnQinAMD Jun 28, 2026
6e59a6d
feat(deepseek-v4): fuse the HCA/CSA compressor forward burst
JohnQinAMD Jun 28, 2026
c91a22b
feat(turbo): fuse expert grouped wgrad -> main_grad (beta=1 ACCUMULAT…
JohnQinAMD Jun 28, 2026
c257479
fix(deepseek-v4): forward PRIMUS_TURBO_FUSE_GROUPED_WGRAD into the co…
JohnQinAMD Jun 28, 2026
67a0539
perf(deepseek-v4): architecture-aware tuned defaults for V4 Triton at…
JohnQinAMD Jun 28, 2026
5eea1cb
perf(deepseek-v4): extend arch-aware V4 attn fwd default to the HCA path
JohnQinAMD Jun 29, 2026
24e344e
doc(deepseek-v4): note V4 attn bwd dKV HG default is MQA-only (inert …
JohnQinAMD Jun 29, 2026
3ce6ee4
dsv4 compressed attn launch reduction (#825)
JohnQinAMD Jun 29, 2026
e32f580
fix: fix padding bug and remove use_turbo_permute_padding flag
RuibinCheung Jun 29, 2026
366beb7
chore: remove unused deepseek-v4 launch scripts
wenxie-amd Jun 30, 2026
c069126
feat(deepseek-v4 projection): manual per-layer timing mode
lhzhang333 Jun 30, 2026
f2f2173
feat: update primus turbo wrapper
RuibinCheung Jul 1, 2026
f179a57
feat: add moe_router_padding_for_quantization flag
RuibinCheung Jul 1, 2026
26bd1cc
FlyDSL + Gluon + bugfix + benchmark + refactor (#831)
wenxie-amd Jul 1, 2026
0fe97fd
feat(projection): add DeepSeek-V4 iteration-timeline view + fix fwd/b…
lhzhang333 Jul 2, 2026
1ab1f8e
style: apply pre-commit lint fixes across primus, deepseek-v4, and tests
wenxie-amd Jul 2, 2026
89bae70
perf(v4-attn): optimize `triton_v2` sparse-MLA forward + backward (#838)
wenxie-amd Jul 2, 2026
46509e7
feat(v4-attn): native FlyDSL sparse-MLA backend (flydsl_v1) + triton_…
wenxie-amd Jul 2, 2026
b3d7669
perf(v4): fuse model-body small kernels into Triton (RMSNorm / HC col…
wenxie-amd Jul 3, 2026
9166dd6
[Megatron-LM] fix: duplicated memory footprint when enable turbo grou…
RuibinCheung Jul 6, 2026
1c9a7d9
feat: add clamped swiglu fusion
RuibinCheung Jul 6, 2026
adf496d
feat(v4-attn): gluon_v2 sparse-MLA backend (Gluon fwd + Gluon bwd) (#…
wenxie-amd Jul 6, 2026
538eba4
opt: remove grouped mlp d2h sync
RuibinCheung Jul 7, 2026
4048787
fix: typo and add PRIMUS_BIAS_SWIGLU_FUSION flag for mi455
RuibinCheung Jul 8, 2026
950833e
ci: update Primus-Turbo/AITER pins, disable torch/jax unittests
wenxie-amd Jul 8, 2026
90eebaf
ci: build and install Triton from source in Docker image
wenxie-amd Jul 8, 2026
75cab09
ci: sync ci.yaml + Dockerfiles from main; keep unittest jobs disabled
wenxie-amd Jul 8, 2026
869ea51
opt: add DeepSeek-V4 gluon v3 attention backend (#865)
wenxie-amd Jul 8, 2026
61df353
feat: use triton to replace jit.fuser to avoid torch.compile bug
RuibinCheung Jul 9, 2026
3fd33f1
feat: remove extra d2h and add swiglu fusion for shared experts
RuibinCheung Jul 10, 2026
0ebfa8c
feat: add triton elementwise add kernel for wgrad accumulate
RuibinCheung Jul 10, 2026
e1a541d
feat(deepseek-v4): add Primus-Turbo native-FlyDSL sparse-MLA attentio…
wenxie-amd Jul 14, 2026
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
768 changes: 529 additions & 239 deletions .github/workflows/ci.yaml

Large diffs are not rendered by default.

2 changes: 1 addition & 1 deletion .github/workflows/deploy-projection.yml
Original file line number Diff line number Diff line change
Expand Up @@ -12,7 +12,7 @@ on:
workflow_dispatch:
push:
branches:
- dev/wenx/deepseek-v4
- dev/tas/deepseek-v4
paths:
- "deepseek-v4/projection/site/**"
- ".github/workflows/deploy-projection.yml"
Expand Down
39 changes: 35 additions & 4 deletions .github/workflows/docker/Dockerfile
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
ARG BASE_IMAGE=docker.io/rocm/primus:v26.2
ARG BASE_IMAGE=docker.io/rocm/primus:v26.3
FROM ${BASE_IMAGE}

ARG PRIMUS_TURBO_COMMIT
Expand All @@ -23,7 +23,7 @@ RUN rm -rf /var/lib/apt/lists/*
ENV ROCSHMEM_HOME=/opt/rocshmem
ENV UCX_HOME=/opt/ucx
# ENV MPI_HOME=/opt/ompi
# Use the system OpenMPI prefix from the v26.2 base image.
# Use the system OpenMPI prefix from the v26.3 base image.
ENV MPI_HOME=/usr/lib/x86_64-linux-gnu/openmpi
ENV ROCM_HOME=/opt/rocm
ENV PRIMUS_TURBO_FRAMEWORK=${PRIMUS_TURBO_FRAMEWORK}
Expand Down Expand Up @@ -55,9 +55,10 @@ RUN rm -rf /opt/rocSHMEM

RUN pip3 uninstall aiter amd-aiter -y || true && \
cd /opt && \
git clone --recursive https://github.com/ROCm/aiter.git && \
git clone https://github.com/ROCm/aiter.git && \
cd aiter && \
git checkout ${PRIMUS_TURBO_AITER_COMMIT} && \
git checkout -f ${PRIMUS_TURBO_AITER_COMMIT} && \
git submodule sync && \
git submodule update --init --recursive && \
PREBUILD_KERNELS=3 pip install --no-cache-dir --use-pep517 .

Expand All @@ -71,6 +72,18 @@ RUN cd /opt && \

RUN rm -rf /opt/Primus-Turbo

# ---------------------------------------------------------------------------
# Install Triton
# ---------------------------------------------------------------------------
RUN cd /opt && \
git clone https://github.com/triton-lang/triton.git && \
cd triton && \
git checkout 09500db9 && \
pip3 install -r python/requirements.txt && \
MAX_JOBS=96 pip3 install --no-build-isolation --force-reinstall --no-deps .

RUN rm -rf /opt/triton

# ---------------------------------------------------------------------------
# Install UCCL-EP (skip for JAX framework)
# ---------------------------------------------------------------------------
Expand All @@ -86,6 +99,24 @@ RUN if [ "$PRIMUS_TURBO_FRAMEWORK" != "JAX" ]; then \
rm -rf /opt/uccl; \
fi

# ---------------------------------------------------------------------------
# Install fixed origami (rocm-libraries@223648a) over the base image's bundled
# 0.1.0. The bundled origami's rank_configs() raises
# `ValueError: vector::reserve` during MoE grouped-gemm kernel selection and
# crashes training (turbo's _safe_rank_configs only catches RuntimeError).
# Skipped for JAX. TODO: drop once a base image ships origami with the fix.
RUN if [ "$PRIMUS_TURBO_FRAMEWORK" != "JAX" ]; then \
rm -rf /tmp/rocm-libraries && \
git clone --filter=blob:none --no-checkout https://github.com/ROCm/rocm-libraries.git /tmp/rocm-libraries && \
cd /tmp/rocm-libraries && \
git sparse-checkout init --cone && \
git sparse-checkout set shared/origami && \
git checkout 223648a26928ebed7f3dd0ccdc044c09f1dccf9b && \
(pip uninstall -y origami || true) && \
pip install --no-cache-dir ./shared/origami/python && \
rm -rf /tmp/rocm-libraries; \
fi

# Set the default working directory
WORKDIR /opt

Expand Down
24 changes: 16 additions & 8 deletions .github/workflows/docker/Dockerfile.ainic
Original file line number Diff line number Diff line change
Expand Up @@ -19,9 +19,9 @@ RUN apt-get update && \
# ---------------------------------------------------------------------------
# Enviroment variables
# ---------------------------------------------------------------------------
ENV WORKDIR=/opt
ENV WORKDIR=/workspace
ENV ROCM_PATH=/opt/rocm
ENV MPI_PATH=/opt/ompi
ENV MPI_PATH=/usr/lib/x86_64-linux-gnu/openmpi

# =============================== Build AINIC Driver ===============================
# WARNING: Please ensure the following environment variables are correctly set:
Expand All @@ -30,16 +30,24 @@ ENV MPI_PATH=/opt/ompi
# WARNING: If these paths are missing, tools and libraries may not function correctly.
# INFO: Installation completed successfully

COPY ${AINIC_BUNDLE_PATH}/ainic_bundle_1.117.5-a-56.tar.gz ${WORKDIR}
COPY ${AINIC_BUNDLE_PATH}/ainic_bundle_1.117.5-a-77.tar.gz ${WORKDIR}
RUN cd ${WORKDIR} && \
rm -rf rccl amd-anp && \
echo "Building ainic bundle... current directory: ${WORKDIR}" && \
tar zxf ainic_bundle_1.117.5-a-56.tar.gz && \
cd ainic_bundle_1.117.5-a-56 && \
tar zxf ainic_bundle_1.117.5-a-77.tar.gz && \
cd ainic_bundle_1.117.5-a-77 && \
tar zxf host_sw_pkg.tar.gz && \
cd host_sw_pkg && \
./install.sh --domain=user -y 2>&1 | tee log_install.txt && \
cd ${WORKDIR} && \
apt-get install -y ./amd/ainic/deb-repo/libionic*.deb
printf '%s\n' \
'Package: libionic1 libionic-dev' \
'Pin: version 54.0-187-1' \
'Pin-Priority: 1001' \
'' \
'Package: perftest' \
'Pin: version 1:25.04.0.0.84-128-1' \
'Pin-Priority: 1001' \
> /etc/apt/preferences.d/ainic-a77 && \
./install.sh --domain=user -y 2>&1 | tee log_install.txt

# ---------------------------------------------------------------------------
# Build rccl
Expand Down
5 changes: 5 additions & 0 deletions .gitignore
Original file line number Diff line number Diff line change
Expand Up @@ -17,3 +17,8 @@ pp_simulation_result
# GitHub Pages) despite the generic `data` ignore above.
!deepseek-v4/projection/site/data/
!deepseek-v4/projection/site/data/*.json

*.log
*.nohup
*.zip
.triton_cache_shared/
112 changes: 112 additions & 0 deletions check_hc_expand.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,112 @@
#!/usr/bin/env python3
"""Correctness + microbench for the fused HC expand Triton kernel vs eager."""
import importlib.util
import os
import time

import torch

os.environ.setdefault("PRIMUS_HC_EXPAND_TRITON", "1")
# Load the kernel module directly by path to avoid the heavy primus/__init__ chain.
_MOD = os.path.join(
os.path.dirname(os.path.abspath(__file__)),
"primus/backends/megatron/core/transformer/v4_attention_kernels/_triton/hc_expand.py",
)
_spec = importlib.util.spec_from_file_location("hc_expand", _MOD)
_hc = importlib.util.module_from_spec(_spec)
_spec.loader.exec_module(_hc)
hc_expand_triton = _hc.hc_expand_triton


def eager_expand(x, out, post, comb):
write = post.unsqueeze(-1) * out.unsqueeze(-2)
mix = torch.matmul(comb, x)
return write + mix


def run(leading, K, D, dtype):
dev = "cuda"
torch.manual_seed(0)
mk = lambda *s: torch.randn(*s, device=dev, dtype=dtype) # noqa: E731
x = mk(*leading, K, D)
out = mk(*leading, D)
post = mk(*leading, K)
comb = torch.softmax(mk(*leading, K, K).float(), dim=-1).to(dtype)

# fp32 reference (ground truth), plus eager and triton at the test dtype.
x32, o32, p32, c32 = (t.float().clone().requires_grad_(True) for t in (x, out, post, comb))
xa, oa, pa, ca = (t.clone().requires_grad_(True) for t in (x, out, post, comb))
xb, ob, pb, cb = (t.clone().requires_grad_(True) for t in (x, out, post, comb))

ref = eager_expand(x32, o32, p32, c32)
eag = eager_expand(xa, oa, pa, ca)
tri = hc_expand_triton(xb, ob, pb, cb)

g32 = torch.randn_like(ref)
g = g32.to(dtype)
ref.backward(g32)
eag.backward(g)
tri.backward(g)

def maxerr(a, b):
return (a.float() - b.float()).abs().max().item()

tag = f"{tuple(leading)} K={K} D={D} {str(dtype).split('.')[-1]}"
ok = True
for name, r, e, t in [
("fwd", ref, eag, tri),
("dx", x32.grad, xa.grad, xb.grad),
("dout", o32.grad, oa.grad, ob.grad),
("dpost", p32.grad, pa.grad, pb.grad),
("dcomb", c32.grad, ca.grad, cb.grad),
]:
e_err = maxerr(r, e) # eager-vs-fp32 (the dtype's intrinsic noise floor)
t_err = maxerr(r, t) # triton-vs-fp32
# Pass if triton is no worse than eager, with an fp32 reduction-order
# slack relative to the tensor magnitude (length-D sums reorder).
slack = 1e-4 + 5e-5 * r.abs().max().item()
passed = t_err <= max(1.5 * e_err, slack)
ok = ok and passed
print(f" {name:6s} triton_err={t_err:.3e} eager_err={e_err:.3e} {'OK' if passed else 'FAIL'}")
print(f"[{tag}] {'PASS' if ok else 'FAIL'}")
return ok


def bench(leading, K, D, dtype, iters=50):
dev = "cuda"
x = torch.randn(*leading, K, D, device=dev, dtype=dtype, requires_grad=True)
out = torch.randn(*leading, D, device=dev, dtype=dtype, requires_grad=True)
post = torch.randn(*leading, K, device=dev, dtype=dtype, requires_grad=True)
comb = (
torch.softmax(torch.randn(*leading, K, K, device=dev).float(), dim=-1).to(dtype).requires_grad_(True)
)
g = torch.randn(*leading, K, D, device=dev, dtype=dtype)

def step(fn):
for t in (x, out, post, comb):
t.grad = None
y = fn(x, out, post, comb)
y.backward(g)

for fn, name in [(eager_expand, "eager"), (hc_expand_triton, "triton")]:
step(fn)
torch.cuda.synchronize()
t0 = time.perf_counter()
for _ in range(iters):
step(fn)
torch.cuda.synchronize()
ms = (time.perf_counter() - t0) / iters * 1e3
print(f" {name:7s} fwd+bwd {ms:8.3f} ms/iter")


if __name__ == "__main__":
print("=== correctness ===")
all_ok = True
all_ok &= run((4096,), 4, 7168, torch.bfloat16)
all_ok &= run((4096,), 4, 7168, torch.float32)
all_ok &= run((2, 17), 4, 320, torch.bfloat16)
all_ok &= run((3,), 2, 128, torch.float32)
all_ok &= run((5,), 8, 64, torch.bfloat16)
print("\n=== bench (production shape B*S=4096, K=4, D=7168, bf16) ===")
bench((4096,), 4, 7168, torch.bfloat16)
print("\nALL", "PASS" if all_ok else "FAIL")
Loading
Loading