[fix] baseline5 fused_add_rms_norm 确实是参数数量不匹配。rms_norm 则是参数数量对了但结果没写回去。两个 bug 叠加导致模型输出全损坏

This commit is contained in:
root
2026-08-19 09:57:51 +00:00
parent 0de522fb4b
commit 9a52068533

View File

@@ -93,14 +93,18 @@ def apply(strict=True):
# vllm uses ops.rms_norm / ops.fused_add_rms_norm
import vllm._custom_ops as ops
def patched_rms_norm(output, input, weight, epsilon):
# vllm: ops.rms_norm(output, input, weight, eps)
# xllm_norm.so: rms_norm(output, input, weight, eps) — same
xllm_ops._get("xllm_norm").rms_norm(output, input, weight, epsilon)
# Cache the C++ module once (avoids dict lookup per call)
_norm_mod = xllm_ops._get("xllm_norm")
def patched_fused_add_rms_norm(input, residual, weight, epsilon):
# vllm: ops.fused_add_rms_norm(input, residual, weight, eps)
# xllm_norm.so: same signature, in-place
def patched_rms_norm(output, input, weight, epsilon):
# C++ signature: rms_norm(output, input, weight, eps)
# Write directly into caller's output tensor (zero-copy).
_norm_mod.rms_norm(output, input, weight, epsilon)
def patched_fused_add_rms_norm(input, residual, weight, epsilon,
residual_alpha=1.0):
if residual_alpha != 1.0:
residual.mul_(residual_alpha)
xllm_ops.residual_rms_norm(input, residual, weight, epsilon)
if hasattr(ops, 'rms_norm'):
@@ -126,9 +130,7 @@ def apply(strict=True):
import vllm._custom_ops as ops
def patched_silu_and_mul(output, input):
# vllm: ops.silu_and_mul(output, input)
# xllm_activation.so: silu_and_mul(out, input) — same order
xllm_ops._get("xllm_activation").silu_and_mul(output, input)
xllm_ops.silu_and_mul(input, output)
if hasattr(ops, 'silu_and_mul'):
ops.silu_and_mul = patched_silu_and_mul
@@ -170,8 +172,7 @@ def apply(strict=True):
import vllm._custom_ops as ops
def patched_reshape_and_cache(key, value, key_cache, value_cache,
slot_mapping, kv_cache_dtype,
k_scale, v_scale):
slot_mapping, kv_cache_dtype, kv_scale):
xllm_ops.reshape_and_cache(key, value, key_cache, value_cache,
slot_mapping)