From 45e17872e2f54e9f5c9d8feec188ff52bd72140b Mon Sep 17 00:00:00 2001 From: Digant Desai Date: Tue, 21 Jul 2026 20:38:04 -0700 Subject: [PATCH 1/2] Update [ghstack-poisoned] --- extension/llm/custom_ops/BUCK | 19 + .../llm/custom_ops/test_quantized_moe.py | 505 ++++++++++++++++++ 2 files changed, 524 insertions(+) create mode 100644 extension/llm/custom_ops/test_quantized_moe.py diff --git a/extension/llm/custom_ops/BUCK b/extension/llm/custom_ops/BUCK index d70b985136d..4b08a639c42 100644 --- a/extension/llm/custom_ops/BUCK +++ b/extension/llm/custom_ops/BUCK @@ -91,3 +91,22 @@ fbcode_target(_kind = runtime.python_test, "//caffe2:torch", ], ) + +fbcode_target(_kind = runtime.python_test, + name = "test_quantized_moe", + srcs = [ + "test_quantized_moe.py", + ], + preload_deps = [ + ":custom_ops_aot_lib_mkl_noomp", + ":custom_ops_aot_py", + "//pytorch/ao/torchao/csrc/cpu/shared_kernels/linear_8bit_act_xbit_weight:op_linear_8bit_act_xbit_weight_aten", + ], + deps = [ + "//caffe2:torch", + "//executorch/extension/pybindings:portable_lib", + "//executorch/examples/models/llama:llama_transformer", + "//executorch/examples/models/llama:transformer_modules", + "//executorch/examples/models/llama:source_transformation", + ], +) diff --git a/extension/llm/custom_ops/test_quantized_moe.py b/extension/llm/custom_ops/test_quantized_moe.py new file mode 100644 index 00000000000..130fe8d331f --- /dev/null +++ b/extension/llm/custom_ops/test_quantized_moe.py @@ -0,0 +1,505 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +# pyre-unsafe + +""" +Python tests for the `llama::quantized_moe_ffn` custom op and the +`replace_moe_with_quantized_op` source transform. + +Numerical correctness is checked by comparing the custom op against a +pure-Python q-dq reference that mirrors `MOEFeedForward.forward` +with the same INT4 group quantization applied to each expert. +""" + +from __future__ import annotations + +import copy +import unittest + +import torch + +from executorch.examples.models.llama.llama_transformer import MOEFeedForward +from executorch.examples.models.llama.model_args import ModelArgs +from executorch.examples.models.llama.source_transformation.moe import ( + _symmetric_quantize_per_group, + QuantizedMoEFFN, + replace_moe_with_quantized_op, +) + +# Importing custom_ops registers the schema + Meta kernel and loads the +# AOT shared library so `torch.ops.llama.quantized_moe_ffn` is callable +# in eager mode. +from executorch.extension.llm.custom_ops import custom_ops # noqa: F401 +from torchao.quantization.granularity import PerGroup +from torchao.quantization.quant_api import ( + Int8DynamicActivationIntxWeightConfig, + quantize_, +) +from torchao.utils import unwrap_tensor_subclass + + +def _qdq_int4_reference(w: torch.Tensor, group_size: int) -> torch.Tensor: + """Apply round-trip symmetric INT4 group quantization. Returns the + "fake-quantized" weight in fp32, matching the values the torchao + kernel would consume.""" + qvals, scales = _symmetric_quantize_per_group(w, group_size) + n, k = w.shape + qvals_grouped = qvals.float().reshape(n, k // group_size, group_size) + scales_grouped = scales.reshape(n, k // group_size, 1) + return (qvals_grouped * scales_grouped).reshape(n, k) + + +def _build_moe_eager( + *, + dim: int = 32, + hidden_dim: int = 32, + num_experts: int = 4, + num_activated_experts: int = 2, + score_func: str = "sigmoid", + use_expert_bias: bool = True, + route_scale: float = 2.5, +) -> MOEFeedForward: + args = ModelArgs( + dim=dim, + n_layers=1, + n_heads=1, + vocab_size=8, + hidden_dim=hidden_dim, + moe=True, + num_experts=num_experts, + num_activated_experts=num_activated_experts, + ) + moe = MOEFeedForward(args) + # `replace_moe_with_quantized_op` reads routing/bias configuration off the + # eager module. Upstream `MOEFeedForward` does not carry these attributes, + # so set them explicitly to drive the transform with the desired config. + moe.num_activated_experts = num_activated_experts + moe.score_func = score_func + moe.route_scale = route_scale + moe.use_expert_bias = use_expert_bias + if use_expert_bias: + moe.expert_bias = torch.randn(num_experts) + moe.eval() + return moe + + +def _moe_forward_with_qdq_weights( + moe: MOEFeedForward, x: torch.Tensor, group_size: int +) -> torch.Tensor: + """Run `moe.forward(x)` after replacing each per-expert weight with its + INT4 group-quantize/dequant round-trip. This is the apples-to-apples + reference for the custom op's output. + + We deepcopy the real module (so any attributes added by future + `MOEFeedForward.__init__` / `ConditionalFeedForward.__init__` changes + are preserved) and only overwrite the per-expert weights in-place. + """ + moe_q = copy.deepcopy(moe) + + def _qdq_per_expert(w_eFD: torch.Tensor) -> torch.Tensor: + # w shape: [E, F, D] for w1/w3 or [E, D, F] for w2 (after the + # transpose in the caller). Quantize each [N, K] slice along K. + return torch.stack( + [_qdq_int4_reference(w_eFD[ei], group_size) for ei in range(w_eFD.size(0))] + ) + + cond_q = moe_q.cond_ffn + with torch.no_grad(): + cond_q.w1.copy_(_qdq_per_expert(moe.cond_ffn.w1)) + cond_q.w3.copy_(_qdq_per_expert(moe.cond_ffn.w3)) + # w2's einsum convention treats it as [F, D] with K=F at packing time. + # The torchao packer is fed `w2.transpose(-2, -1)` (shape [E, D, F]), + # quantized along K=F, then dequantized. Equivalent to quantizing each + # [D, F] slice and transposing back to [F, D]. + w2_DF = moe.cond_ffn.w2.transpose(-2, -1).contiguous() + w2_DF_qdq = _qdq_per_expert(w2_DF) + cond_q.w2.copy_(w2_DF_qdq.transpose(-2, -1).contiguous()) + + return moe_q.forward(x) + + +class TestSourceTransform(unittest.TestCase): + """Tests for `replace_moe_with_quantized_op`.""" + + def test_replaces_module_in_place_with_correct_buffers(self) -> None: + moe = _build_moe_eager( + dim=32, hidden_dim=32, num_experts=4, num_activated_experts=2 + ) + wrapper = torch.nn.Module() + wrapper.block_sparse_moe = moe + replace_moe_with_quantized_op(wrapper, group_size=32, weight_nbit=4) + replaced = wrapper.block_sparse_moe + self.assertIsInstance(replaced, QuantizedMoEFFN) + self.assertEqual(replaced.num_experts, 4) + self.assertEqual(replaced.num_activated_experts, 2) + self.assertEqual(replaced.dim, 32) + self.assertEqual(replaced.hidden_dim, 32) + self.assertEqual(tuple(replaced.gate_weight.shape), (4, 32)) + self.assertEqual(replaced.expert_bias.numel(), 4) + self.assertEqual(replaced.packed_w1.shape[0], 4) + self.assertEqual(replaced.packed_w3.shape[0], 4) + self.assertEqual(replaced.packed_w2.shape[0], 4) + self.assertEqual(replaced.packed_w1.shape[1], replaced.packed_w3.shape[1]) + + def test_torchao_quantized_gate_exports_as_fp32_buffer(self) -> None: + moe = _build_moe_eager( + dim=32, hidden_dim=32, num_experts=4, num_activated_experts=2 + ) + config = Int8DynamicActivationIntxWeightConfig( + weight_dtype=torch.int4, + weight_granularity=PerGroup(32), + weight_scale_dtype=torch.bfloat16, + ) + quantize_(moe, config, filter_fn=lambda _module, fqn: fqn == "gate") + moe = unwrap_tensor_subclass(moe) + + wrapper = torch.nn.Module() + wrapper.block_sparse_moe = moe + replace_moe_with_quantized_op(wrapper, group_size=32, weight_nbit=4) + + replaced = wrapper.block_sparse_moe + self.assertEqual(replaced.gate_weight.dtype, torch.float32) + exported = torch.export.export(replaced, (torch.randn(2, 32),)) + call_targets = { + node.target for node in exported.graph.nodes if node.op == "call_function" + } + self.assertIn(torch.ops.llama.quantized_moe_ffn.default, call_targets) + + def test_no_expert_bias_produces_empty_buffer(self) -> None: + moe = _build_moe_eager( + dim=32, + hidden_dim=32, + num_experts=4, + num_activated_experts=2, + use_expert_bias=False, + ) + wrapper = torch.nn.Module() + wrapper.block_sparse_moe = moe + replace_moe_with_quantized_op(wrapper, group_size=32, weight_nbit=4) + replaced = wrapper.block_sparse_moe + self.assertIsInstance(replaced, QuantizedMoEFFN) + self.assertEqual(replaced.expert_bias.numel(), 0) + + def test_buffers_stay_fp32_after_to_half(self) -> None: + moe = _build_moe_eager( + dim=32, hidden_dim=32, num_experts=4, num_activated_experts=2 + ) + wrapper = torch.nn.Module() + wrapper.block_sparse_moe = moe + replace_moe_with_quantized_op(wrapper, group_size=32, weight_nbit=4) + wrapper.to(torch.float16) + replaced = wrapper.block_sparse_moe + self.assertEqual(replaced.gate_weight.dtype, torch.float32) + self.assertEqual(replaced.expert_bias.dtype, torch.float32) + self.assertEqual(replaced.packed_w1.dtype, torch.uint8) + + def test_buffers_stay_fp32_after_to_bfloat16(self) -> None: + moe = _build_moe_eager( + dim=32, hidden_dim=32, num_experts=4, num_activated_experts=2 + ) + wrapper = torch.nn.Module() + wrapper.block_sparse_moe = moe + replace_moe_with_quantized_op(wrapper, group_size=32, weight_nbit=4) + wrapper.to(torch.bfloat16) + replaced = wrapper.block_sparse_moe + self.assertEqual(replaced.gate_weight.dtype, torch.float32) + self.assertEqual(replaced.expert_bias.dtype, torch.float32) + + def test_3d_input_shape_preserved_through_forward(self) -> None: + moe = _build_moe_eager( + dim=32, hidden_dim=32, num_experts=4, num_activated_experts=2 + ) + wrapper = torch.nn.Module() + wrapper.block_sparse_moe = moe + replace_moe_with_quantized_op(wrapper, group_size=32, weight_nbit=4) + replaced = wrapper.block_sparse_moe + x_meta = torch.empty((2, 4, 32), dtype=torch.float32, device="meta") + out = torch.ops.llama.quantized_moe_ffn( + x_meta.view(-1, 32), + replaced.gate_weight.to("meta"), + replaced.expert_bias.to("meta"), + replaced.packed_w1.to("meta"), + replaced.packed_w3.to("meta"), + replaced.packed_w2.to("meta"), + replaced.num_activated_experts, + replaced.num_experts, + replaced.hidden_dim, + replaced.dim, + replaced.group_size, + replaced.weight_nbit, + replaced.score_func, + replaced.route_scale, + ) + self.assertEqual(tuple(out.shape), (8, 32)) + out_reshaped = out.view(2, 4, 32) + self.assertEqual(tuple(out_reshaped.shape), (2, 4, 32)) + + def test_forward_output_dtype_matches_input(self) -> None: + moe = _build_moe_eager( + dim=32, hidden_dim=32, num_experts=4, num_activated_experts=2 + ) + wrapper = torch.nn.Module() + wrapper.block_sparse_moe = moe + replace_moe_with_quantized_op(wrapper, group_size=32, weight_nbit=4) + x_fp16 = torch.randn(4, 32, dtype=torch.float16) + with torch.no_grad(): + out = wrapper.block_sparse_moe(x_fp16) + self.assertEqual(out.dtype, torch.float16) + + def test_nested_replacement(self) -> None: + moe = _build_moe_eager( + dim=32, hidden_dim=32, num_experts=4, num_activated_experts=2 + ) + outer = torch.nn.Module() + inner = torch.nn.Module() + inner.feed_forward = moe + outer.layer0 = inner + replace_moe_with_quantized_op(outer, group_size=32, weight_nbit=4) + self.assertIsInstance(outer.layer0.feed_forward, QuantizedMoEFFN) + + def test_meta_kernel_returns_correct_output_shape(self) -> None: + moe = _build_moe_eager(num_experts=4, num_activated_experts=2) + wrapper = torch.nn.Module() + wrapper.block_sparse_moe = moe + replace_moe_with_quantized_op(wrapper, group_size=32, weight_nbit=4) + replaced = wrapper.block_sparse_moe + x_meta = torch.empty((8, replaced.dim), dtype=torch.float32, device="meta") + out_meta = torch.ops.llama.quantized_moe_ffn( + x_meta, + replaced.gate_weight.to("meta"), + replaced.expert_bias.to("meta"), + replaced.packed_w1.to("meta"), + replaced.packed_w3.to("meta"), + replaced.packed_w2.to("meta"), + replaced.num_activated_experts, + replaced.num_experts, + replaced.hidden_dim, + replaced.dim, + replaced.group_size, + replaced.weight_nbit, + replaced.score_func, + replaced.route_scale, + ) + self.assertEqual(tuple(out_meta.shape), (8, replaced.dim)) + self.assertEqual(out_meta.dtype, torch.float32) + + def test_meta_kernel_single_token(self) -> None: + moe = _build_moe_eager(num_experts=4, num_activated_experts=2) + wrapper = torch.nn.Module() + wrapper.block_sparse_moe = moe + replace_moe_with_quantized_op(wrapper, group_size=32, weight_nbit=4) + replaced = wrapper.block_sparse_moe + x_meta = torch.empty((1, replaced.dim), dtype=torch.float32, device="meta") + out_meta = torch.ops.llama.quantized_moe_ffn( + x_meta, + replaced.gate_weight.to("meta"), + replaced.expert_bias.to("meta"), + replaced.packed_w1.to("meta"), + replaced.packed_w3.to("meta"), + replaced.packed_w2.to("meta"), + replaced.num_activated_experts, + replaced.num_experts, + replaced.hidden_dim, + replaced.dim, + replaced.group_size, + replaced.weight_nbit, + replaced.score_func, + replaced.route_scale, + ) + self.assertEqual(tuple(out_meta.shape), (1, replaced.dim)) + + +class TestInt4PackRoundtrip(unittest.TestCase): + """The torchao INT4 packer + unpacker should preserve weight values + within INT4 quantization resolution.""" + + def test_pack_unpack_within_quant_resolution(self) -> None: + torch.manual_seed(1) + n, k = 32, 64 + w = torch.randn(n, k) + qvals, scales = _symmetric_quantize_per_group(w, group_size=32) + recon = (qvals.float().reshape(n, 2, 32) * scales.reshape(n, 2, 1)).reshape( + n, k + ) + max_step = (w.abs().reshape(n, 2, 32).amax(dim=-1) / 7.0).max().item() + diff = (w - recon).abs().max().item() + self.assertLessEqual(diff, max_step + 1e-6) + + def test_full_row_group_has_lower_error(self) -> None: + torch.manual_seed(2) + n, k = 32, 64 + w = torch.randn(n, k) + _, _ = _symmetric_quantize_per_group(w, group_size=32) + qvals_fine, scales_fine = _symmetric_quantize_per_group(w, group_size=32) + recon_fine = ( + qvals_fine.float().reshape(n, 2, 32) * scales_fine.reshape(n, 2, 1) + ).reshape(n, k) + qvals_coarse, scales_coarse = _symmetric_quantize_per_group(w, group_size=k) + recon_coarse = ( + qvals_coarse.float().reshape(n, 1, k) * scales_coarse.reshape(n, 1, 1) + ).reshape(n, k) + err_fine = (w - recon_fine).abs().max().item() + err_coarse = (w - recon_coarse).abs().max().item() + self.assertGreaterEqual(err_coarse, err_fine * 0.5) + + +class TestMetaKernelValidation(unittest.TestCase): + """Meta kernel input validation rejects bad inputs.""" + + def _make_valid_inputs(self) -> dict: + moe = _build_moe_eager( + dim=32, hidden_dim=32, num_experts=4, num_activated_experts=2 + ) + wrapper = torch.nn.Module() + wrapper.m = moe + replace_moe_with_quantized_op(wrapper, group_size=32, weight_nbit=4) + r = wrapper.m + return dict( + x=torch.empty((4, r.dim), dtype=torch.float32, device="meta"), + gate_weight=r.gate_weight.to("meta"), + expert_bias=r.expert_bias.to("meta"), + packed_w1=r.packed_w1.to("meta"), + packed_w3=r.packed_w3.to("meta"), + packed_w2=r.packed_w2.to("meta"), + num_activated_experts=r.num_activated_experts, + num_experts=r.num_experts, + hidden_dim=r.hidden_dim, + dim=r.dim, + group_size=r.group_size, + weight_nbit=r.weight_nbit, + score_func=r.score_func, + route_scale=r.route_scale, + ) + + def test_rejects_3d_input(self) -> None: + kw = self._make_valid_inputs() + kw["x"] = torch.empty((1, 4, 32), dtype=torch.float32, device="meta") + with self.assertRaises(AssertionError): + torch.ops.llama.quantized_moe_ffn(**kw) + + def test_rejects_bad_score_func(self) -> None: + kw = self._make_valid_inputs() + kw["score_func"] = "gelu" + with self.assertRaises(AssertionError): + torch.ops.llama.quantized_moe_ffn(**kw) + + def test_rejects_bad_weight_nbit(self) -> None: + kw = self._make_valid_inputs() + kw["weight_nbit"] = 3 + with self.assertRaises(AssertionError): + torch.ops.llama.quantized_moe_ffn(**kw) + + def test_rejects_num_activated_experts_zero(self) -> None: + kw = self._make_valid_inputs() + kw["num_activated_experts"] = 0 + with self.assertRaises(AssertionError): + torch.ops.llama.quantized_moe_ffn(**kw) + + def test_rejects_non_uint8_packed_w1(self) -> None: + kw = self._make_valid_inputs() + kw["packed_w1"] = kw["packed_w1"].to(torch.float32) + with self.assertRaises(AssertionError): + torch.ops.llama.quantized_moe_ffn(**kw) + + def test_rejects_non_uint8_packed_w2(self) -> None: + kw = self._make_valid_inputs() + kw["packed_w2"] = kw["packed_w2"].to(torch.float32) + with self.assertRaises(AssertionError): + torch.ops.llama.quantized_moe_ffn(**kw) + + def test_rejects_non_uint8_packed_w3(self) -> None: + kw = self._make_valid_inputs() + kw["packed_w3"] = kw["packed_w3"].to(torch.float32) + with self.assertRaises(AssertionError): + torch.ops.llama.quantized_moe_ffn(**kw) + + def test_rejects_mismatched_packed_w1_w3_sizes(self) -> None: + kw = self._make_valid_inputs() + e = kw["num_experts"] + kw["packed_w3"] = torch.empty( + (e, kw["packed_w3"].size(1) + 1), dtype=torch.uint8, device="meta" + ) + with self.assertRaises(AssertionError): + torch.ops.llama.quantized_moe_ffn(**kw) + + def test_rejects_non_fp32_gate_weight(self) -> None: + kw = self._make_valid_inputs() + kw["gate_weight"] = kw["gate_weight"].to(torch.float16) + with self.assertRaises(AssertionError): + torch.ops.llama.quantized_moe_ffn(**kw) + + def test_rejects_non_fp32_expert_bias(self) -> None: + kw = self._make_valid_inputs() + kw["expert_bias"] = kw["expert_bias"].to(torch.float16) + with self.assertRaises(AssertionError): + torch.ops.llama.quantized_moe_ffn(**kw) + + +class TestKernelAvailability(unittest.TestCase): + """The kernel (reference or optimized) is always available.""" + + def test_sentinel_op_registered(self) -> None: + self.assertTrue(hasattr(torch.ops.llama, "_quantized_moe_ffn_active")) + self.assertTrue(torch.ops.llama._quantized_moe_ffn_active()) + + +class TestGroupSizeValidation(unittest.TestCase): + """Validate group_size and divisibility checks.""" + + def test_rejects_group_size_zero(self) -> None: + kw = TestMetaKernelValidation()._make_valid_inputs() + kw["group_size"] = 0 + with self.assertRaises(AssertionError): + torch.ops.llama.quantized_moe_ffn(**kw) + + def test_rejects_dim_not_divisible_by_group_size(self) -> None: + kw = TestMetaKernelValidation()._make_valid_inputs() + kw["group_size"] = 7 + with self.assertRaises(AssertionError): + torch.ops.llama.quantized_moe_ffn(**kw) + + def test_source_transform_rejects_invalid_weight_nbit(self) -> None: + moe = _build_moe_eager( + dim=32, hidden_dim=32, num_experts=4, num_activated_experts=2 + ) + wrapper = torch.nn.Module() + wrapper.m = moe + with self.assertRaises(ValueError): + replace_moe_with_quantized_op(wrapper, group_size=32, weight_nbit=3) + + def test_source_transform_rejects_group_size_zero(self) -> None: + moe = _build_moe_eager( + dim=32, hidden_dim=32, num_experts=4, num_activated_experts=2 + ) + wrapper = torch.nn.Module() + wrapper.m = moe + with self.assertRaises(ValueError): + replace_moe_with_quantized_op(wrapper, group_size=0, weight_nbit=4) + + +class TestSharedExpert(unittest.TestCase): + """Shared expert is preserved through the source transform.""" + + def test_shared_expert_preserved(self) -> None: + moe = _build_moe_eager( + dim=32, hidden_dim=32, num_experts=4, num_activated_experts=2 + ) + moe.shared_expert = torch.nn.Linear(32, 32) + wrapper = torch.nn.Module() + wrapper.m = moe + replace_moe_with_quantized_op(wrapper, group_size=32, weight_nbit=4) + self.assertIsNotNone(wrapper.m.shared_expert) + self.assertIsInstance(wrapper.m.shared_expert, torch.nn.Linear) + + def test_no_shared_expert_is_none(self) -> None: + moe = _build_moe_eager( + dim=32, hidden_dim=32, num_experts=4, num_activated_experts=2 + ) + wrapper = torch.nn.Module() + wrapper.m = moe + replace_moe_with_quantized_op(wrapper, group_size=32, weight_nbit=4) + self.assertIsNone(wrapper.m.shared_expert) From 6f7bf1336cdfb8051bca63a71f5619433a3cd7f7 Mon Sep 17 00:00:00 2001 From: Digant Desai Date: Wed, 22 Jul 2026 13:33:23 -0700 Subject: [PATCH 2/2] Update [ghstack-poisoned] --- .../llm/custom_ops/test_quantized_moe.py | 32 +++++++++---------- 1 file changed, 16 insertions(+), 16 deletions(-) diff --git a/extension/llm/custom_ops/test_quantized_moe.py b/extension/llm/custom_ops/test_quantized_moe.py index 130fe8d331f..47ba58207f3 100644 --- a/extension/llm/custom_ops/test_quantized_moe.py +++ b/extension/llm/custom_ops/test_quantized_moe.py @@ -358,22 +358,22 @@ def _make_valid_inputs(self) -> dict: wrapper.m = moe replace_moe_with_quantized_op(wrapper, group_size=32, weight_nbit=4) r = wrapper.m - return dict( - x=torch.empty((4, r.dim), dtype=torch.float32, device="meta"), - gate_weight=r.gate_weight.to("meta"), - expert_bias=r.expert_bias.to("meta"), - packed_w1=r.packed_w1.to("meta"), - packed_w3=r.packed_w3.to("meta"), - packed_w2=r.packed_w2.to("meta"), - num_activated_experts=r.num_activated_experts, - num_experts=r.num_experts, - hidden_dim=r.hidden_dim, - dim=r.dim, - group_size=r.group_size, - weight_nbit=r.weight_nbit, - score_func=r.score_func, - route_scale=r.route_scale, - ) + return { + "x": torch.empty((4, r.dim), dtype=torch.float32, device="meta"), + "gate_weight": r.gate_weight.to("meta"), + "expert_bias": r.expert_bias.to("meta"), + "packed_w1": r.packed_w1.to("meta"), + "packed_w3": r.packed_w3.to("meta"), + "packed_w2": r.packed_w2.to("meta"), + "num_activated_experts": r.num_activated_experts, + "num_experts": r.num_experts, + "hidden_dim": r.hidden_dim, + "dim": r.dim, + "group_size": r.group_size, + "weight_nbit": r.weight_nbit, + "score_func": r.score_func, + "route_scale": r.route_scale, + } def test_rejects_3d_input(self) -> None: kw = self._make_valid_inputs()