Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
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
Original file line number Diff line number Diff line change
Expand Up @@ -18,7 +18,7 @@ This guide uses Slurm and the `trtllm-llmapi-launch` multi-node launcher. The co
| TEP8 | attention-TP, EP8 | 213 GB | GB300-class per-GPU memory |
| TEP16 | attention-TP, EP16 | 115 GB | validated on GB200 (`SM100`) |

DEP16 replicates the BF16 non-expert weights on every rank (114 GB per rank) on top of the MXFP4 routed experts at 16-way expert parallelism (90 GB per rank). TEP16 shards those non-expert weights instead, which is what brings it within `SM100` per-GPU memory; TEP8 does not fit because its 8-way expert share alone is 181 GB per rank. On B200 (`SM100`), Kimi K3 is functionally supported at the kernel and module level and covered by unit tests in CI. Other GPU architectures are not supported.
DEP16 replicates the BF16 non-expert weights on every rank (114 GB per rank) on top of the MXFP4 routed experts at 16-way expert parallelism (90 GB per rank). TEP16 shards those non-expert weights instead, which is what brings it within `SM100` per-GPU memory; TEP8 does not fit because its 8-way expert share alone is 181 GB per rank. On B200 (`SM100`), Kimi K3 is functionally supported at the kernel and module level and covered by unit tests in CI. Full-model deployments on other GPU architectures are not covered by this guide.
* Multi-node launcher: Slurm with the pyxis/enroot container plugin (or an equivalent MPI launcher) to start one rank per GPU across the nodes.
* High-speed inter-node interconnect (e.g., NVLink/InfiniBand) for the expert-parallel traffic.
* Shared filesystem visible to all nodes for the repository, the model weights, and the configuration file.
Expand All @@ -34,7 +34,8 @@ The checkpoint and the configuration file must live on a shared filesystem visib

## Feature Support Notes

* **Blackwell only.** NVIDIA Blackwell GPUs are supported. The performance results in this guide were validated on NVIDIA GB300 NVL GPUs. Kimi K3 kernels and modules are also functional on B200 (`SM100`) and covered by unit tests in CI, and the TEP16 deployment is validated end-to-end on GB200 (`SM100`); DEP16 and TEP8 require GB300-class per-GPU memory (see Prerequisites). Support for other GPU architectures may be added in a future release.
* **Full-model deployments in this guide use Blackwell.** The performance results were validated on NVIDIA GB300 NVL GPUs. Kimi K3 kernels and modules are also functional on B200 (`SM100`) and covered by unit tests in CI, and the TEP16 deployment is validated end-to-end on GB200 (`SM100`); DEP16 and TEP8 require GB300-class per-GPU memory (see Prerequisites).
* **Rubin FC12 MoE kernel support.** On SM107, `CUTEDSL_FC12` supports SwiGLU and Kimi K3 SiTU for NVFP4 routed experts. This kernel/backend support does not validate the full Kimi K3 deployment on Rubin; the configurations and results below are for Blackwell.
* **High-throughput and low-latency deployments are provided.** DEP16 (`enable_attention_dp: true`, `moe_expert_parallel_size: 16`) is the high-throughput deployment. TEP16 (`enable_attention_dp: false`, `moe_expert_parallel_size: 16`) is the low-latency deployment. An 8-GPU deployment, TEP8 (`enable_attention_dp: false`, `moe_expert_parallel_size: 8`), is also provided. Select the deployment and concurrency appropriate for your workload.
* **CUDA graphs and the overlap scheduler are enabled.** The performance-sweep recipes set `disable_overlap_scheduler: false` and enable CUDA graphs. DEP16 additionally sets `cuda_graph_config.enable_padding: true`.
* **Chunked prefill is supported and enabled** (`enable_chunked_prefill: true`), so prompts longer than `max_num_tokens` are scheduled across multiple steps.
Expand Down
57 changes: 43 additions & 14 deletions tensorrt_llm/_torch/custom_ops/cute_dsl_custom_ops.py
Original file line number Diff line number Diff line change
Expand Up @@ -17287,7 +17287,7 @@ def _(
# builds that do not ship cutlass.utils.rubin_helpers (SM107 only).
if IS_CUTLASS_DSL_RUBIN_AVAILABLE:
# ----------------------------------------------------------------
# Rubin NVFP4 Fused FC12 (FC1 gather+SwiGLU + FC2 finalize in ONE kernel)
# Rubin NVFP4 Fused FC12 (FC1 gather+gated-act + FC2 finalize in ONE kernel)
# ----------------------------------------------------------------
# Compat shim: the delivered fused kernel accesses ``cutlass.memory.*``
# (SmemAllocator / TmemAllocator / get_smem_capacity_in_bytes) and
Expand Down Expand Up @@ -17380,16 +17380,20 @@ class Sm107BlockScaledContiguousGroupedGemmFusedFc12Runner(
kernel_cache = dict()
tuning_config_cache = dict()

def __init__(self,
num_experts: int,
top_k: int,
num_local_experts: int,
local_expert_offset: int,
tile_size: int,
scaling_vector_size: int = 16,
swiglu_limit: float = float("inf"),
ep_size: int = 1,
enable_alltoall: bool = False):
def __init__(
self,
num_experts: int,
top_k: int,
num_local_experts: int,
local_expert_offset: int,
tile_size: int,
scaling_vector_size: int = 16,
swiglu_limit: float = float("inf"),
ep_size: int = 1,
enable_alltoall: bool = False,
activation_type: ActivationType = ActivationType.Swiglu,
situ_beta: Optional[float] = None,
situ_linear_beta: Optional[float] = None):
super().__init__()
self.num_experts = num_experts
self.top_k = top_k
Expand All @@ -17398,6 +17402,9 @@ def __init__(self,
self.tile_size = tile_size
self.scaling_vector_size = scaling_vector_size
self.swiglu_limit = swiglu_limit
self.activation_type = ActivationType(activation_type)
self.situ_beta = situ_beta
self.situ_linear_beta = situ_linear_beta
# Used only by the in-op output memset (moved here so the memset
# is the fused kernel's immediate stream predecessor).
self.ep_size = ep_size
Expand All @@ -17422,6 +17429,9 @@ def unique_id(self):
self.tile_size,
self.scaling_vector_size,
self.swiglu_limit,
int(self.activation_type),
self.situ_beta,
self.situ_linear_beta,
)

def get_valid_tactics(
Expand Down Expand Up @@ -17719,7 +17729,8 @@ def forward(self, inputs: List[torch.Tensor],
cache_key = (self.scaling_vector_size, self.tile_size,
self.top_k, mma_tiler, mma_inst_shape,
cluster_shape_mn, max_active_clusters,
self.swiglu_limit)
self.swiglu_limit, int(self.activation_type),
self.situ_beta, self.situ_linear_beta)
if cache_key not in self.__class__.kernel_cache:
gemm = self.__class__.kernel_class(
self.scaling_vector_size,
Expand All @@ -17731,6 +17742,9 @@ def forward(self, inputs: List[torch.Tensor],
use_pdl=True,
swiglu_limit=self.swiglu_limit,
scheduler="l2_atomic",
activation_type=self.activation_type,
situ_beta=self.situ_beta,
situ_linear_beta=self.situ_linear_beta,
)

@cute.jit
Expand Down Expand Up @@ -17966,6 +17980,9 @@ def _run_nvfp4_fc12_fused_rubin(
ep_size: int,
enable_alltoall: bool,
tuner_key: str,
activation_type: ActivationType = ActivationType.Swiglu,
situ_beta: Optional[float] = None,
situ_linear_beta: Optional[float] = None,
) -> torch.Tensor:
tuner = AutoTuner.get()
runner = Sm107BlockScaledContiguousGroupedGemmFusedFc12Runner(
Expand All @@ -17978,6 +17995,9 @@ def _run_nvfp4_fc12_fused_rubin(
swiglu_limit=swiglu_limit,
ep_size=ep_size,
enable_alltoall=enable_alltoall,
activation_type=activation_type,
situ_beta=situ_beta,
situ_linear_beta=situ_linear_beta,
)
# Input order matches Fc12FusedInputsHelper (FC1 prefix 0..9 mirrors
# GatherGroupedGemmInputsHelper; FC2 tensors 10..14; expanded_idx 15
Expand Down Expand Up @@ -18014,7 +18034,9 @@ def _run_nvfp4_fc12_fused_rubin(
"SymInt local_expert_offset, SymInt tile_size, float swiglu_limit, "
"SymInt ep_size, bool enable_alltoall, "
"SymInt scaling_vector_size=16, "
"str? precomputed_tactic=None) -> ()",
"str? precomputed_tactic=None, "
f"SymInt activation_type={int(ActivationType.Swiglu)}, "
"float? situ_beta=None, float? situ_linear_beta=None) -> ()",
device_types="cuda")
def cute_dsl_nvfp4_fc12_fused_rubin(
input: torch.Tensor,
Expand Down Expand Up @@ -18043,6 +18065,9 @@ def cute_dsl_nvfp4_fc12_fused_rubin(
enable_alltoall: bool,
scaling_vector_size: int = 16,
precomputed_tactic: Optional[str] = None,
activation_type: int = int(ActivationType.Swiglu),
situ_beta: Optional[float] = None,
situ_linear_beta: Optional[float] = None,
) -> None:
# In-place: finalize scatter-adds into ``output`` (mutates_args);
# the op returns nothing so it does not alias its own input.
Expand All @@ -18055,7 +18080,8 @@ def cute_dsl_nvfp4_fc12_fused_rubin(
local_expert_offset, tile_size, scaling_vector_size,
swiglu_limit, precomputed_tactic, expanded_idx_to_permuted_idx,
ep_size, enable_alltoall,
"trtllm::cute_dsl_nvfp4_fc12_fused_rubin")
"trtllm::cute_dsl_nvfp4_fc12_fused_rubin",
ActivationType(activation_type), situ_beta, situ_linear_beta)

@torch.library.register_fake("trtllm::cute_dsl_nvfp4_fc12_fused_rubin")
def _(
Expand Down Expand Up @@ -18085,5 +18111,8 @@ def _(
enable_alltoall: bool,
scaling_vector_size: int = 16,
precomputed_tactic: Optional[str] = None,
activation_type: int = int(ActivationType.Swiglu),
situ_beta: Optional[float] = None,
situ_linear_beta: Optional[float] = None,
) -> None:
return None
Original file line number Diff line number Diff line change
Expand Up @@ -2857,6 +2857,9 @@ def _apply_situ_epilogue(

beta * tanh(x/beta) = beta * (2*sigmoid(2x/beta) - 1)
= 2*beta*sigmoid((2/beta)*x) - beta

Keep the packed SiTU algebra in sync with the Rubin fused FC12
``rubin_contiguous_grouped_blockscaled_gemm_fused_fc12.py`` epilogue.
"""
beta = self.situ_beta
linear_beta = self.situ_linear_beta
Expand Down
Loading
Loading