From 65a4fc4e2dec5ea9f2ecbd963a4e7c9acbbca80a Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Wed, 19 Mar 2025 02:39:50 -0700 Subject: [PATCH] Update vision.py --- unsloth/models/vision.py | 21 +++++++++++++-------- 1 file changed, 13 insertions(+), 8 deletions(-) diff --git a/unsloth/models/vision.py b/unsloth/models/vision.py index c5cd57de97..0497ce4379 100644 --- a/unsloth/models/vision.py +++ b/unsloth/models/vision.py @@ -177,25 +177,30 @@ def unsloth_base_fast_generate( if "attention_mask" in kwargs: torch._dynamo.mark_static(kwargs["attention_mask"], 0) torch._dynamo.mark_dynamic(kwargs["attention_mask"], 1) - if "pixel_values" in kwargs: - print(kwargs["pixel_values"].shape) if "token_type_ids" in kwargs: torch._dynamo.mark_static(kwargs["token_type_ids"], 0) torch._dynamo.mark_dynamic(kwargs["token_type_ids"], 1) # Fix generation_config - cache_implementation = getattr(self.config, "cache_implementation", "static") + # Use hybrid if sliding window seen, otherwise try static + cache_implementation = getattr(self.config, "cache_implementation", None) + if cache_implementation is None: + swa = getattr(getattr(model.config, "text_config", model.config), "sliding_window", None) + if swa == 0 or type(swa) is not int: + cache_implementation = "static" + else: + cache_implementation = "hybrid" + if getattr(self, "_supports_static_cache", True): + cache_implementation = "static" + else: + cache_implementation = None if "generation_config" in kwargs: kwargs["generation_config"].cache_implementation = cache_implementation kwargs["generation_config"].compile_config = _compile_config - elif getattr(self, "_supports_static_cache", True): + else: kwargs["cache_implementation"] = cache_implementation kwargs["compile_config"] = _compile_config - else: - kwargs["cache_implementation"] = "hybrid" - kwargs["compile_config"] = _compile_config pass - print(kwargs) with torch.inference_mode(), autocaster: try: