diff --git a/python/sglang/multimodal_gen/runtime/managers/cpu_worker.py b/python/sglang/multimodal_gen/runtime/managers/cpu_worker.py index e596665b1..69cf9af23 100644 --- a/python/sglang/multimodal_gen/runtime/managers/cpu_worker.py +++ b/python/sglang/multimodal_gen/runtime/managers/cpu_worker.py @@ -74,7 +74,3 @@ class CPUWorker(GPUWorker): # Set local size to hint SGLang to use shared memory based AllReduce os.environ["LOCAL_SIZE"] = str(self.server_args.tp_size) torch.ops.sgl_kernel.initialize(self.server_args.tp_size, self.rank) - - @torch.library.register_fake("sgl_kernel::shm_allgather") - def _(data, dim): - return torch.cat([data] * self.server_args.tp_size, dim=dim) diff --git a/python/sglang/srt/model_executor/cpu_graph_runner.py b/python/sglang/srt/model_executor/cpu_graph_runner.py index 8c059f7f4..c8e5a199f 100644 --- a/python/sglang/srt/model_executor/cpu_graph_runner.py +++ b/python/sglang/srt/model_executor/cpu_graph_runner.py @@ -119,6 +119,7 @@ def set_torch_compile_config(): torch._dynamo.config.accumulated_cache_size_limit = 1024 if hasattr(torch._dynamo.config, "cache_size_limit"): torch._dynamo.config.cache_size_limit = 1024 + register_inductor_fallback_ops() monkey_patch_torch_compile() @@ -138,7 +139,28 @@ def get_batch_sizes_to_capture(model_runner: ModelRunner): return capture_bs -def register_fake_ops(): +_CPU_COMPILE_FAKE_OPS: set[str] = set() + + +def register_cpu_compile_fake(op_name: str): + _CPU_COMPILE_FAKE_OPS.add(op_name) + return torch.library.register_fake(f"sgl_kernel::{op_name}") + + +def register_inductor_fallback_ops(): + from torch._inductor.lowering import lowerings, make_fallback + + sgl_kernel_ops = torch.ops.sgl_kernel + for op_name in sorted(_CPU_COMPILE_FAKE_OPS): + try: + op = getattr(getattr(sgl_kernel_ops, op_name), "default") + except AttributeError: + continue + if op not in lowerings: + make_fallback(op, warn=False) + + +def register_fake_ops(tp_size: int): """ Registers fake/meta implementations for all custom sgl_kernel CPU operators using torch.library.register_fake to support torch.compile @@ -156,7 +178,7 @@ def register_fake_ops(): ] for op in none_return_ops: - @torch.library.register_fake(f"sgl_kernel::{op}") + @register_cpu_compile_fake(op) def _(*args, **kwargs): return @@ -173,11 +195,15 @@ def register_fake_ops(): "gemma4_rmsnorm_cpu", ]: - @torch.library.register_fake(f"sgl_kernel::{op}") + @register_cpu_compile_fake(op) def _(input, *args, **kwargs): return torch.empty_like(input) - @torch.library.register_fake("sgl_kernel::qkv_proj_with_rope") + @register_cpu_compile_fake("shm_allgather") + def _(data, dim): + return torch.cat([data] * tp_size, dim=dim) + + @register_cpu_compile_fake("qkv_proj_with_rope") def _( hidden_states, q_a_proj_weight, @@ -218,14 +244,18 @@ def register_fake_ops(): v_input = k_input.narrow(-1, 0, kv_lora_rank) return q_input, k_input, v_input - @torch.library.register_fake("sgl_kernel::rotary_embedding_cpu") + @register_cpu_compile_fake("rotary_embedding_cpu") def _(positions, query, key, head_size, cos_sin_cache, is_neox): if query.ndim == 2: return query, key else: return torch.empty_like(query), torch.empty_like(key) - @torch.library.register_fake("sgl_kernel::multimodal_rotary_embedding_cpu") + @register_cpu_compile_fake("apply_rotary_pos_emb_cpu") + def _(query, key, cos, sin): + return query, key + + @register_cpu_compile_fake("multimodal_rotary_embedding_cpu") def _( positions, query, @@ -238,7 +268,7 @@ def register_fake_ops(): ): return query, key - @torch.library.register_fake("sgl_kernel::qkv_proj_with_rope_fused_weight") + @register_cpu_compile_fake("qkv_proj_with_rope_fused_weight") def _( hidden_states, q_a_proj_weight, @@ -292,13 +322,13 @@ def register_fake_ops(): return mat2.shape[1] return mat2.shape[0] - @torch.library.register_fake("sgl_kernel::weight_packed_linear") + @register_cpu_compile_fake("weight_packed_linear") def _(mat1, mat2, bias, is_vnni): M = mat1.shape[0] N = get_n_size(mat2, is_vnni) return mat1.new_empty(M, N) - @torch.library.register_fake("sgl_kernel::per_token_quant_int8_cpu") + @register_cpu_compile_fake("per_token_quant_int8_cpu") def _(input): M = input.shape[0] K = input.shape[1] @@ -306,14 +336,14 @@ def register_fake_ops(): As = input.new_empty(M, dtype=torch.float32) return Aq, As - @torch.library.register_fake("sgl_kernel::int8_scaled_mm_cpu") + @register_cpu_compile_fake("int8_scaled_mm_cpu") def _(mat1, mat2, scales1, scales2, bias, out_dtype, is_vnni): M = mat1.shape[0] N = mat2.shape[0] out = mat1.new_empty(M, N, dtype=out_dtype) return out - @torch.library.register_fake("sgl_kernel::grouped_topk_cpu") + @register_cpu_compile_fake("grouped_topk_cpu") def _( hidden_states, gating_output, @@ -332,7 +362,7 @@ def register_fake_ops(): topk_ids = torch.empty(shape, device=device, dtype=torch.int) return topk_weights, topk_ids - @torch.library.register_fake("sgl_kernel::biased_grouped_topk_cpu") + @register_cpu_compile_fake("biased_grouped_topk_cpu") def _( hidden_states, gating_output, @@ -352,7 +382,7 @@ def register_fake_ops(): topk_ids = torch.empty(shape, device=device, dtype=torch.int) return topk_weights, topk_ids - @torch.library.register_fake("sgl_kernel::topk_sigmoid_cpu") + @register_cpu_compile_fake("topk_sigmoid_cpu") def _(hidden_states, gating_output, topk, renormalize): num_tokens = hidden_states.shape[0] shape = (num_tokens, topk) @@ -361,7 +391,7 @@ def register_fake_ops(): torch.empty(shape, device=hidden_states.device, dtype=torch.int), ) - @torch.library.register_fake("sgl_kernel::topk_softmax_cpu") + @register_cpu_compile_fake("topk_softmax_cpu") def _( hidden_states, gating_output, @@ -381,7 +411,7 @@ def register_fake_ops(): "gelu_and_mul_cpu", ]: - @torch.library.register_fake(f"sgl_kernel::{act_op}") + @register_cpu_compile_fake(act_op) def _(input): sizes = list(input.shape) last_dim = input.dim() - 1 @@ -389,7 +419,7 @@ def register_fake_ops(): sizes[last_dim] = d return input.new_empty(sizes) - @torch.library.register_fake("sgl_kernel::int8_scaled_mm_with_quant") + @register_cpu_compile_fake("int8_scaled_mm_with_quant") def _( mat1, mat2, @@ -402,7 +432,7 @@ def register_fake_ops(): N = mat2.shape[0] return mat1.new_empty(M, N, dtype=out_dtype) - @torch.library.register_fake("sgl_kernel::fp8_scaled_mm_cpu") + @register_cpu_compile_fake("fp8_scaled_mm_cpu") def _( mat1, mat2, @@ -416,7 +446,19 @@ def register_fake_ops(): N = mat2.shape[0] return mat1.new_empty(M, N, dtype=out_dtype) - @torch.library.register_fake("sgl_kernel::fused_linear_sigmoid_mul") + @register_cpu_compile_fake("mxfp4_scaled_mm_cpu") + def _(mat1, mat2, scales2, bias, is_vnni): + sizes = list(mat1.shape) + sizes[-1] = mat2.shape[0] + return mat1.new_empty(sizes) + + @register_cpu_compile_fake("int4_scaled_mm_cpu") + def _(x, w, w_zeros, w_scales, bias): + sizes = list(x.shape) + sizes[-1] = w_scales.shape[0] * w_scales.shape[-1] + return x.new_empty(sizes) + + @register_cpu_compile_fake("fused_linear_sigmoid_mul") def _( mat1, mat2, @@ -428,7 +470,7 @@ def register_fake_ops(): N = post_mul_mat.shape[1] return mat1.new_empty(M, N) - @torch.library.register_fake("sgl_kernel::fused_qkvzba_split_reshape_cat_cpu") + @register_cpu_compile_fake("fused_qkvzba_split_reshape_cat_cpu") def _(mixed_qkvz, mixed_ba, num_heads_qk, num_heads_v, head_qk, head_v): batch = mixed_qkvz.shape[0] qkv_dim = num_heads_qk * head_qk * 2 + num_heads_v * head_v @@ -438,9 +480,7 @@ def register_fake_ops(): a = mixed_ba.new_empty(batch, num_heads_v) return mixed_qkv, z, b, a - @torch.library.register_fake( - "sgl_kernel::fused_qkvzba_split_reshape_cat_contiguous_cpu" - ) + @register_cpu_compile_fake("fused_qkvzba_split_reshape_cat_contiguous_cpu") def _(mixed_qkvz, mixed_ba, num_heads_qk, num_heads_v, head_qk, head_v): batch = mixed_qkvz.shape[0] qkv_dim = num_heads_qk * head_qk * 2 + num_heads_v * head_v @@ -450,9 +490,7 @@ def register_fake_ops(): a = mixed_ba.new_empty(batch, num_heads_v) return mixed_qkv, z, b, a - @torch.library.register_fake( - "sgl_kernel::fused_sigmoid_gating_delta_rule_update_cpu" - ) + @register_cpu_compile_fake("fused_sigmoid_gating_delta_rule_update_cpu") def _( A_log, dt_bias, @@ -476,7 +514,7 @@ def register_fake_ops(): v_head_dim = v.shape[3] return q.new_empty(batch_size, seq_len, v_num_heads, v_head_dim) - @torch.library.register_fake("sgl_kernel::fused_gdn_gating_cpu") + @register_cpu_compile_fake("fused_gdn_gating_cpu") def _(A_log, a, b, dt_bias): batch = a.shape[0] num_heads = a.shape[1] @@ -484,7 +522,7 @@ def register_fake_ops(): beta = b.new_empty(1, batch, num_heads) return out, beta - @torch.library.register_fake("sgl_kernel::chunk_gated_delta_rule_cpu") + @register_cpu_compile_fake("chunk_gated_delta_rule_cpu") def _( query, key, @@ -587,7 +625,7 @@ class CPUGraphRunner: ) if self.enable_torch_compile: - register_fake_ops() + register_fake_ops(self.tp_size) set_torch_compile_config() # Graph inputs diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index 42b428830..e91075fa7 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -1148,10 +1148,6 @@ class ModelRunner(ModelRunnerKVCacheMixin): os.environ["LOCAL_SIZE"] = str(self.tp_size) torch.ops.sgl_kernel.initialize(self.tp_size, self.tp_rank) - @torch.library.register_fake("sgl_kernel::shm_allgather") - def _(data, dim): - return torch.cat([data] * self.tp_size, dim=dim) - else: logger.warning( "init_cpu_threads_env and shared memory based AllReduce is disabled, only intel amx backend and arm64 are supported"