From dc42a0d82ce066d18a5f6f7e8a9caaeced0633b7 Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Sun, 17 May 2026 14:55:03 +0000 Subject: [PATCH] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- unsloth/models/gemma4_moe_4bit.py | 11 ++++++----- 1 file changed, 6 insertions(+), 5 deletions(-) diff --git a/unsloth/models/gemma4_moe_4bit.py b/unsloth/models/gemma4_moe_4bit.py index 13727de986..6d85412ffc 100644 --- a/unsloth/models/gemma4_moe_4bit.py +++ b/unsloth/models/gemma4_moe_4bit.py @@ -174,6 +174,7 @@ def _ensure_pt_dequant_state(layer): if getattr(layer, "_unsloth_pt_dequant_ready", False): return from bitsandbytes.functional import dequantize_blockwise + qs = layer.weight.quant_state if qs.nested: absmax_fp32 = dequantize_blockwise(qs.absmax, qs.state2) @@ -199,7 +200,7 @@ def _pt_dequant_one(packed_uint8, absmax_fp32, blocksize, shape, dtype, codes): n_blocks = (n_elements + blocksize - 1) // blocksize values = values.view(n_blocks, blocksize) * absmax_fp32.view(-1, 1) target = shape[0] * shape[1] - return values.reshape(-1)[: target].view(shape).to(dtype) + return values.reshape(-1)[:target].view(shape).to(dtype) def _pt_dequant_stack_subset(layers, indices_cpu, codes): @@ -232,9 +233,6 @@ def _get_compiled_pt_dequant_stack(): return _COMPILED_PT_DEQUANT_STACK - - - def _dequant_stack(layers): """Dequantize each Linear4bit in a ModuleList and stack into (E, out, in).""" from bitsandbytes.functional import dequantize_4bit @@ -413,7 +411,10 @@ def _grouped_mm_forward_4bit_pt_compiled( self.down_proj = nn.Parameter(down, requires_grad = False) try: return forward_native_grouped_mm( - self, hidden_states, compact_top_k, top_k_weights, + self, + hidden_states, + compact_top_k, + top_k_weights, ) finally: del self.gate_up_proj