From 8ebf07c9b1727490c66d8dbc3195cd24d7b98870 Mon Sep 17 00:00:00 2001 From: Zhenshan Xie Date: Sun, 20 Sep 2026 00:45:15 -0700 Subject: [PATCH 1/2] feat(qwen3_8): support native BF16 precision Qwen3.8-27B's own checkpoint is natively BF16; qwen3_8 previously only accepted precision=fp16 or fp32, forcing a downcast even when a user wants to match the reference checkpoint's own numeric format. Wires precision="bf16" through model.py's validation and engine_builder.py's work_np_dtype/work_trt_dtype selection. Following families/qwen's existing, proven BF16 pattern: weight constants are staged as FP16 bytes (TensorRT's Weights constructor does not accept ml_dtypes.bfloat16 arrays directly) and the network's runtime dtype is trt.bfloat16; every generic constant-building helper in graph_ops.py already casts its constant to match its activation's runtime dtype via _cast_back_to_trt_dtype, so this propagates automatically through most of the graph. Two DeltaNet-specific call sites built raw elementwise ops directly against a freshly-created constant without going through that cast-matching helper, which is fine when the constant and activation happen to already share a dtype (the previously-only-supported fp16/fp32 modes) but breaks once the constant is staged as FP16 while the activation is genuinely BF16: the conv1d weight multiply in the DeltaNet conv step, and the GQA head-tiling "ones" constant used to expand compact Q/K from num_kv_heads to num_heads. Both now explicitly cast the constant to the activation's dtype before the elementwise op. Verified: builds successfully against the original Qwen/Qwen3.8-27B BF16 checkpoint (precision="bf16"), produces coherent generation output, and IEngineInspector confirms 273 GEMM layers use real BF16 tensor-core tactics (sm80_xmma_gemm_bf16bf16_..., cutlass3x_sm100_tensorop..._bf16...), not a fallback. DeltaNet's recurrent state update correctly stays in FP32 throughout, matching the reference HF implementation's own numerical behavior (unchanged from the existing fp16/fp32 modes). Signed-off-by: Zhenshan Xie --- families/qwen3_8/engine_builder.py | 16 ++++++++++++++-- families/qwen3_8/model.py | 4 ++-- 2 files changed, 16 insertions(+), 4 deletions(-) diff --git a/families/qwen3_8/engine_builder.py b/families/qwen3_8/engine_builder.py index d43d795c4..67a486e91 100644 --- a/families/qwen3_8/engine_builder.py +++ b/families/qwen3_8/engine_builder.py @@ -481,11 +481,19 @@ def build_engine( partial_rotary_factor: float = weights["_partial_rotary_factor"] if precision == "fp16": work_np_dtype, work_trt_dtype = np.float16, trt.float16 + elif precision == "bf16": + # Constants are staged as FP16 bytes (TensorRT's Weights constructor + # does not accept ml_dtypes.bfloat16 arrays directly) and explicitly + # cast to BF16 in-graph by graph_ops._cast_back_to_trt_dtype, which + # every constant-building helper already calls to match its + # activation's runtime dtype -- mirroring families/qwen's own + # "storage np_dtype is fp16, runtime trt_dtype is bfloat16" pattern. + work_np_dtype, work_trt_dtype = np.float16, trt.bfloat16 elif precision == "fp32": work_np_dtype, work_trt_dtype = np.float32, trt.float32 else: raise ValueError( - f"Unsupported Qwen3.8 precision {precision!r}; expected fp32 or fp16") + f"Unsupported Qwen3.8 precision {precision!r}; expected fp32, fp16, or bf16") requested_fp32_layers = frozenset( int(layer) for layer in config.raw.get("_fp32_layers", ())) invalid_fp32_layers = sorted( @@ -603,7 +611,7 @@ def build_engine( prefix = f"layer.{layer_idx}" lt = layer_types[layer_idx] layer_is_fp32 = ( - precision == "fp16" and layer_idx in requested_fp32_layers) + precision in ("fp16", "bf16") and layer_idx in requested_fp32_layers) layer_np_dtype = np.float32 if layer_is_fp32 else work_np_dtype layer_trt_dtype = trt.float32 if layer_is_fp32 else work_trt_dtype @@ -897,6 +905,8 @@ def _add_deltanet_layer( conv_w = graph_ops.add_constant( network, (conv_dim, d_conv), weights[f"{prefix}.conv1d_weight"], dtype=dtype) + if conv_w.dtype != present_conv.dtype: + conv_w = network.add_cast(conv_w, present_conv.dtype).get_output(0) conv_prod = network.add_elementwise( present_conv, conv_w, trt.ElementWiseOperation.PROD) conv_sum = network.add_reduce( @@ -950,6 +960,8 @@ def _add_deltanet_layer( tile_ones = graph_ops.add_constant( network, (1, heads_per_group, 1), np.ones((1, heads_per_group, 1), dtype=dtype), dtype=dtype) + if tile_ones.dtype != q_3d.get_output(0).dtype: + tile_ones = network.add_cast(tile_ones, q_3d.get_output(0).dtype).get_output(0) q_tiled = network.add_elementwise( q_3d.get_output(0), tile_ones, trt.ElementWiseOperation.PROD) q_expanded_s = network.add_shuffle(q_tiled.get_output(0)) diff --git a/families/qwen3_8/model.py b/families/qwen3_8/model.py index 1897c6cab..b35ad73f5 100644 --- a/families/qwen3_8/model.py +++ b/families/qwen3_8/model.py @@ -110,8 +110,8 @@ def build(request, writer) -> None: ): raise ValueError("checkpoint is not a Qwen3.8 model") precision = str(request.precision).lower() - if precision not in {"fp16", "fp32"}: - raise ValueError("qwen3_8 precision must be fp16 or fp32") + if precision not in {"fp16", "fp32", "bf16"}: + raise ValueError("qwen3_8 precision must be fp16, bf16, or fp32") max_sequence_length = _positive_int( request.max_sequence_length or min(config.max_position_embeddings, 256), "max_sequence_length", From b190b13f753866169afd78b556aabdef4a9a34fc Mon Sep 17 00:00:00 2001 From: Zhenshan Xie Date: Sun, 20 Sep 2026 15:02:19 -0700 Subject: [PATCH 2/2] feat(qwen3_8): default to BF16 precision when unspecified Qwen3.8-27B's own checkpoint is natively BF16; the family's default_precision was still "fp32" (the base FamilySupport default), so omitting --precision silently built at a precision the checkpoint was never published in. Set default_precision="bf16" on both the alias-matched and config-discriminated FamilySupport instances, so `trtmc build -o out.bundle` (no --precision flag) now matches the checkpoint's own native format by default. Explicit --precision fp16/fp32 continue to work unchanged. Signed-off-by: Zhenshan Xie --- families/qwen3_8/support.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/families/qwen3_8/support.py b/families/qwen3_8/support.py index 015d90530..6db925d74 100644 --- a/families/qwen3_8/support.py +++ b/families/qwen3_8/support.py @@ -10,8 +10,10 @@ model_types=("qwen38", "qwen3.8", "qwen3_8"), tasks=("text_generation",), default_task="text_generation", + default_precision="bf16", ) -_SUPPORT = FamilySupport(tasks=("text_generation",), default_task="text_generation") +_SUPPORT = FamilySupport( + tasks=("text_generation",), default_task="text_generation", default_precision="bf16") def describe(metadata: ModelMetadata) -> FamilySupport | None: