Skip to content
Merged
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
50 changes: 50 additions & 0 deletions csrc/jit_kernels/impls/runtime_utils.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -4,13 +4,63 @@
#include <torch/python.h>

#include "../heuristics/sm90.hpp"
#include "../../jit/device_runtime.hpp"
#include "../../jit/handle.hpp"
#include "../../utils/math.hpp"
#include "../../utils/system.hpp"
#include "../../utils/exception.hpp"

namespace deep_gemm {

// Grid size for the cooperative mega-MoE kernels, with SM headroom reserved for
// co-resident kernels.
//
// These kernels launch a persistent grid of one CTA per SM joined by a software
// whole-grid barrier (`comm/barrier.cuh`) with a hard 60s timeout. The barrier is
// only satisfiable if every CTA is *simultaneously resident*, but an ordinary
// launch gives no gang-scheduling guarantee: CUDA places CTAs greedily, so a
// concurrent kernel holding some (but not all) SMs leaves part of the grid queued
// while the resident CTAs spin against the deadline and never yield their slots.
// The result is a mutual deadlock that traps at 60s -> 'DeepGEMM grid sync
// timeout' (barrier.cuh:39) / 'NVLink barrier timeout' (barrier.cuh:80), which in
// turn produces Xid 43 and kills the process.
//
// Reserving SMs gives a later-arriving co-resident kernel somewhere to land other
// than an SM the grid needs. How much headroom is enough depends entirely on the
// competing kernel's footprint, which is deployment-specific:
// - disaggregated serving: NCCL `ncclDevKernel_SendRecv` (KV transfer) — 2 was
// measured sufficient (yingz/mega-sm-headroom)
// - training: FSDP2 all-gather / reduce-scatter / HSDP all-reduce on dedicated
// comm streams; a 32-channel collective touches ~8 SMs (one channel = one
// 512-thread block, 4 blocks/SM), so 8 (with NCCL_MAX_NCHANNELS=8 pinned)
// was measured for the INC-1291 shape
//
// There is deliberately NO built-in default: the headroom is set exclusively via
// `DG_MEGA_MOE_SM_HEADROOM`, and an unset variable means 0 (no reservation), i.e.
// the historical full-device grid. Deployments that run comm kernels concurrently
// with the mega grid MUST set it; see the INC-1291 RCA (fw-ai/fireworks#47752)
// for how to size it.
//
// The value is clamped even, because these kernels launch 2-CTA clusters (2-SM
// MMA with a paired TMEM accumulator layout) and several call sites assert
// `num_sms % 2 == 0`.
static int get_mega_moe_num_sms() {
const int num_device_sms = device_runtime->get_num_sms();
// No default: unset env == 0 == the historical full-device grid. The
// deployment owns the value (see comment above).
int headroom = get_env<int>("DG_MEGA_MOE_SM_HEADROOM", 0);

DG_HOST_ASSERT(headroom >= 0 and "DG_MEGA_MOE_SM_HEADROOM must be non-negative");
// Round up to keep the grid even for the 2-CTA cluster launch.
headroom = align(headroom, 2);
DG_HOST_ASSERT(headroom < num_device_sms and
"DG_MEGA_MOE_SM_HEADROOM leaves no SMs for the mega-MoE grid");

const int num_sms = num_device_sms - headroom;
DG_HOST_ASSERT(num_sms % 2 == 0);
return num_sms;
}

static std::pair<int, int> get_inner_outer_dims(const cute::UMMA::Major& major, const int& k, const int& mn) {
return major == cute::UMMA::Major::K ? std::make_pair(k, mn) : std::make_pair(mn, k);
}
Expand Down
4 changes: 3 additions & 1 deletion csrc/jit_kernels/impls/sm100_bf16_mega_moe.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -296,9 +296,11 @@ static void sm100_bf16_mega_moe(

// Launch
const auto physical_num_sms = device_runtime->get_num_sms();
// Default to the shared headroom-reserving grid size (see
// `get_mega_moe_num_sms`); the explicit absolute override is retained.
const auto num_sms = get_env<int>(
"DG_BF16_MEGA_MOE_NUM_SMS",
physical_num_sms);
get_mega_moe_num_sms());
DG_HOST_ASSERT(num_sms > 0 && num_sms <= physical_num_sms);
const SM100BF16MegaMoERuntime::Args args = {
.num_max_tokens_per_rank = num_max_tokens_per_rank,
Expand Down
2 changes: 1 addition & 1 deletion csrc/jit_kernels/impls/sm100_fp8_fp4_mega_moe.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -326,7 +326,7 @@ static void sm100_fp8_fp4_mega_moe(
cumulative_local_expert_recv_stats_ptr = cumulative_local_expert_recv_stats->data_ptr<int>();

// Launch
const auto num_sms = device_runtime->get_num_sms();
const int num_sms = get_mega_moe_num_sms();
const SM100FP8FP4MegaMoERuntime::Args args = {
.num_max_tokens_per_rank = num_max_tokens_per_rank,
.hidden = hidden, .intermediate_hidden = intermediate_hidden,
Expand Down
8 changes: 4 additions & 4 deletions csrc/jit_kernels/impls/sm100_fp8_fp4_mega_moe_backward.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -743,7 +743,7 @@ static void sm100_fp8_fp4_mega_moe_backward_dgrad_swiglu(
dgrad_block_k, load_block_m,
static_cast<int>(grad_gate_up_output.stride(-2)), 128);
// Each launch gets a unique readiness epoch; no host memset is required.
const int num_sms = device_runtime->get_num_sms();
const int num_sms = get_mega_moe_num_sms();
DG_HOST_ASSERT(num_sms % 2 == 0);
static std::atomic<uint32_t> next_launch_epoch{1};
uint32_t launch_epoch =
Expand Down Expand Up @@ -953,7 +953,7 @@ static void sm100_mega_moe_backward_combine_grad_x(
DG_HOST_ASSERT(topk_ids->size(1) == num_topk);
}

const int num_sms = device_runtime->get_num_sms();
const int num_sms = get_mega_moe_num_sms();
const SM100MegaMoEBackwardCombineRuntime::Args args = {
.num_ranks = num_ranks,
.num_local_experts = num_local_experts,
Expand Down Expand Up @@ -1145,7 +1145,7 @@ static void sm100_bf16_mega_moe_backward_post_down_prelude(
}
}

const int num_sms = device_runtime->get_num_sms();
const int num_sms = get_mega_moe_num_sms();
const auto backward_sym_buffer = layout::SymBuffer<>(
backward_sym_buffer_ptrs, backward_rank);
const auto backward_workspace = layout::Workspace(
Expand Down Expand Up @@ -1605,7 +1605,7 @@ static void sm100_bf16_mega_moe_backward_dgrad(
static_cast<int>(
grad_gate_up_output.stride(-2)), 128);

const int num_sms = device_runtime->get_num_sms();
const int num_sms = get_mega_moe_num_sms();
DG_HOST_ASSERT(num_sms % 2 == 0);
constexpr int num_trace_sites = 22;
constexpr int num_trace_values = 5;
Expand Down