diff --git a/tests/test_generate_kwarg_gate.py b/tests/test_generate_kwarg_gate.py index 242fcebfbf..1073160f29 100644 --- a/tests/test_generate_kwarg_gate.py +++ b/tests/test_generate_kwarg_gate.py @@ -138,6 +138,7 @@ def test_generate_kwarg_gate(): # The only values that must be stripped are the ones the strict validator would # raise on, which is exactly what the gate above predicts. + def _filter_logits_kwargs(model, kwargs): """The v5 branch of unsloth_base_fast_generate, as a testable function.""" for key in ("logits_to_keep", "num_logits_to_keep"): @@ -169,9 +170,10 @@ def test_v5_leaves_other_kwargs_alone(): def test_source_has_no_unconditional_pop(): src = open(VISION).read() - assert 'kwargs.pop("logits_to_keep", None)\n kwargs.pop("num_logits_to_keep", None)' not in src, ( - "the v5 branch must not drop caller-supplied logits_to_keep unconditionally" - ) + assert ( + 'kwargs.pop("logits_to_keep", None)\n kwargs.pop("num_logits_to_keep", None)' + not in src + ), "the v5 branch must not drop caller-supplied logits_to_keep unconditionally" if __name__ == "__main__": diff --git a/unsloth/models/vision.py b/unsloth/models/vision.py index effe633302..08af645e4f 100644 --- a/unsloth/models/vision.py +++ b/unsloth/models/vision.py @@ -431,8 +431,7 @@ def unsloth_base_fast_generate(self, *args, **kwargs): # (renamed away in v5), and logits_to_keep on the VLMs whose top-level # forward does not take it. for _logits_kwarg in ("logits_to_keep", "num_logits_to_keep"): - if _logits_kwarg in kwargs and \ - not _unsloth_generate_accepts_kwarg(self, _logits_kwarg): + if _logits_kwarg in kwargs and not _unsloth_generate_accepts_kwarg(self, _logits_kwarg): kwargs.pop(_logits_kwarg, None) model_eos_token_id = getattr(self.config, "eos_token_id", None)