From e44d3054d5fb2d0f5b607b71a8ea1566c2bc9c20 Mon Sep 17 00:00:00 2001 From: potoior <2986485901@qq.com> Date: Mon, 17 Aug 2026 14:42:25 +0800 Subject: [PATCH] fix(qwen3_5): skip GDN input-projection fusion on mixed projection dtypes Anti-compressed NVFP4 checkpoints can quantize only some GDN input projections (e.g. in_proj_qkv/in_proj_z -> NVFP4 FP8 weights) while leaving others (in_proj_b/in_proj_a) FP16. fuse_gdn_input_projections checked only the first projection's quant type and then concat weights of all four projections, so these checkpoints crashed with 'Promotion for Float8 Types is not supported, attempted to promote Float8_e4m3fn and Half'. Verify all four projections share the same weight dtype before fusing; when they differ, log a warning and keep the unfused path (which is already supported in forward()). Pure FP16 / pure NVFP4 models still fuse. Verified: tensorrt-edgellm-export Qwen3.8-27B-NVFP4 --skip-audio completes, with 'GDN fusion skipped' warnings logged per layer. Signed-off-by: potoior <2986485901@qq.com> --- .../models/qwen3_5/modeling_qwen3_5_text.py | 11 +++++++++++ 1 file changed, 11 insertions(+) diff --git a/tensorrt_edgellm/models/qwen3_5/modeling_qwen3_5_text.py b/tensorrt_edgellm/models/qwen3_5/modeling_qwen3_5_text.py index 47a3b8a24..70b1ebe7a 100644 --- a/tensorrt_edgellm/models/qwen3_5/modeling_qwen3_5_text.py +++ b/tensorrt_edgellm/models/qwen3_5/modeling_qwen3_5_text.py @@ -807,6 +807,17 @@ def fuse_gdn_input_projections(model: nn.Module) -> int: # --- Fuse: concatenate weights along output dim (dim 0) ---------- fused_buffers: dict = {} proj_modules = [getattr(mixer, n) for n in _GDN_PROJ_NAMES] + # Anti-compressed checkpoints can quantize only some GDN projections + # (e.g. in_proj_qkv/z -> NVFP4 FP8, in_proj_b/a -> FP16); a single cat + # cannot mix those dtypes, so keep the unfused path in that case. + wptrs = [getattr(p, "weight", None) for p in proj_modules] + if any(wp is None for wp in wptrs) or len({wp.dtype + for wp in wptrs}) != 1: + logger.warning( + "GDN fusion skipped for %s: projections have mixed weight " + "dtypes (%s).", name, + ", ".join(sorted(str(wp.dtype) for wp in wptrs))) + continue for attr in list(proj_modules[0]._buffers) + list( proj_modules[0]._parameters): parts = [getattr(p, attr) for p in proj_modules]