Improve registration in cpu_graph_runner (#26635)
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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"
|
||||
|
||||
Reference in New Issue
Block a user