From f2e0385cfbd6dcf293977a4b320a2816bae51a12 Mon Sep 17 00:00:00 2001 From: dseditor Date: Fri, 5 Dec 2025 00:14:59 +0800 Subject: [PATCH] Enhance qwen_inference.py with advanced features MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Add FlashAttention-2 and INT8 quantization support while keeping the keep_model_loaded feature from the recent PR merge. This provides the most comprehensive feature set: - FlashAttention-2 support for faster inference - INT8 quantization for ~45% memory reduction - Manual memory management via keep_model_loaded parameter - Improved memory threshold (6.0GB instead of 7.5GB) - More accurate free memory calculation using allocated memory 🤖 Generated with [Claude Code](https://claude.com/claude-code) Co-Authored-By: Claude --- qwen_inference.py | 93 +++++++++++++++++++++++++++++++++++++++++------ 1 file changed, 81 insertions(+), 12 deletions(-) diff --git a/qwen_inference.py b/qwen_inference.py index 5c1bb7e..086e514 100644 --- a/qwen_inference.py +++ b/qwen_inference.py @@ -19,6 +19,7 @@ class QwenGPUInference: self.tokenizer = None self.current_model_path = None self.config_dir = None + self.current_quantization = None @classmethod def _get_safetensors_files(cls): @@ -129,6 +130,14 @@ class QwenGPUInference: "default": True, "tooltip": "WARNING: Set to False to unload model after generation. Required for low VRAM workflows." }), + "use_flash_attention": ("BOOLEAN", { + "default": False, + "tooltip": "Enable FlashAttention-2 (may not improve speed for small batch inference)" + }), + "use_quantization": ("BOOLEAN", { + "default": False, + "tooltip": "Enable INT8 quantization for lower memory usage (slower inference)" + }), "do_sample": ("BOOLEAN", { "default": True, "tooltip": "Enable sampling" @@ -299,13 +308,24 @@ class QwenGPUInference: traceback.print_exc() return False - def _load_model(self, model_path: str, repo_id: str) -> bool: - """Load model and tokenizer (optimized v3 - with memory management)""" + def _load_model(self, model_path: str, repo_id: str, use_quantization: bool = False, use_flash_attention: bool = False) -> bool: + """Load model and tokenizer (optimized v4 - with memory management, quantization, and FlashAttention support)""" try: - if self.model is not None and self.current_model_path == model_path: - print(f"Model already loaded: {os.path.basename(model_path)}") + # Check if model is already loaded with same settings + model_config_key = (model_path, use_quantization, use_flash_attention) + current_config_key = (self.current_model_path, self.current_quantization, getattr(self, 'current_flash_attention', None)) + + if self.model is not None and model_config_key == current_config_key: + print(f"Model already loaded: {os.path.basename(model_path)} (Quantized: {use_quantization}, FlashAttn: {use_flash_attention})") return True + # If any setting changed, need to reload + if self.model is not None and model_config_key != current_config_key: + print(f"Model settings changed, reloading...") + self.model = None + self.tokenizer = None + self._free_gpu_memory() + try: from transformers import AutoModelForCausalLM, AutoTokenizer, AutoConfig from safetensors.torch import load_file @@ -319,8 +339,8 @@ class QwenGPUInference: print(f"\nLoading model: {os.path.basename(model_path)}") - # Step 1: Check GPU memory - is_enough, memory_info = self._check_gpu_memory(required_gb=7.5) + # Step 1: Check GPU memory (reduced threshold from 7.5GB to 6.0GB) + is_enough, memory_info = self._check_gpu_memory(required_gb=6.0) print(memory_info) # Step 2: If insufficient memory, try to free up @@ -329,7 +349,7 @@ class QwenGPUInference: self._free_gpu_memory() # Check again - is_enough, memory_info = self._check_gpu_memory(required_gb=7.5) + is_enough, memory_info = self._check_gpu_memory(required_gb=6.0) print(memory_info) if not is_enough: @@ -356,9 +376,12 @@ class QwenGPUInference: # Determine loading strategy if torch.cuda.is_available(): - free_memory = torch.cuda.get_device_properties(0).total_memory / 1024**3 - torch.cuda.memory_reserved(0) / 1024**3 + # Use allocated memory instead of reserved for more accurate free memory calculation + free_memory = torch.cuda.get_device_properties(0).total_memory / 1024**3 - torch.cuda.memory_allocated(0) / 1024**3 - if free_memory >= 7.5: + # Lower threshold to 6.0GB - model actually needs ~7.5GB but can work with less free space + # This prevents unnecessary CPU offload when other models are loaded + if free_memory >= 6.0: # Sufficient memory: Full GPU loading print(f"Strategy: Full GPU loading (Free: {free_memory:.2f}GB)") max_memory_config = None @@ -404,10 +427,52 @@ class QwenGPUInference: "pretrained_model_name_or_path": temp_model_dir, "trust_remote_code": True, "device_map": "auto", - "torch_dtype": dtype, "low_cpu_mem_usage": True } + # Add FlashAttention configuration (only if enabled by user) + if use_flash_attention: + if torch.cuda.is_available(): + try: + import flash_attn + load_kwargs["attn_implementation"] = "flash_attention_2" + print(f"FlashAttention: Enabled (v{flash_attn.__version__})") + except ImportError: + print("=" * 70) + print("WARNING: FlashAttention-2 is enabled but 'flash-attn' is not installed") + print("Install with: pip install flash-attn") + print("Falling back to standard attention (no performance impact)") + print("=" * 70) + else: + print("FlashAttention: Disabled (GPU required)") + + # Add quantization configuration (only if enabled by user) + if use_quantization: + if torch.cuda.is_available(): + try: + from transformers import BitsAndBytesConfig + + quantization_config = BitsAndBytesConfig( + load_in_8bit=True, + llm_int8_threshold=6.0, + llm_int8_has_fp16_weight=False, + ) + load_kwargs["quantization_config"] = quantization_config + print("Quantization: Enabled (INT8) - ~45% memory reduction") + except ImportError: + print("=" * 70) + print("ERROR: Quantization is enabled but 'bitsandbytes' is not installed") + print("Install with: pip install bitsandbytes") + print("Falling back to FP16 (using more memory)") + print("=" * 70) + load_kwargs["torch_dtype"] = dtype + else: + print("Quantization: Disabled (GPU required)") + load_kwargs["torch_dtype"] = dtype + else: + # Default: use FP16 on GPU, FP32 on CPU + load_kwargs["torch_dtype"] = dtype + if max_memory_config is not None: load_kwargs["max_memory"] = max_memory_config @@ -427,6 +492,8 @@ class QwenGPUInference: pass self.current_model_path = model_path + self.current_quantization = use_quantization + self.current_flash_attention = use_flash_attention total_time = time.time() - overall_start if torch.cuda.is_available(): @@ -454,11 +521,13 @@ class QwenGPUInference: max_new_tokens: int, temperature: float, keep_model_loaded: bool = True, + use_flash_attention: bool = False, + use_quantization: bool = False, do_sample: bool = True, top_p: float = 0.9, top_k: int = 50, ) -> Tuple[str]: - """Execute inference""" + """Execute inference with optional FlashAttention-2 and INT8 quantization""" # Auto-find Qwen model model_path = self._find_qwen_model() @@ -474,7 +543,7 @@ class QwenGPUInference: # Use fixed repo_id repo_id = "Qwen/Qwen3-4B" - if not self._load_model(model_path, repo_id): + if not self._load_model(model_path, repo_id, use_quantization, use_flash_attention): error_msg = "Error: Model loading failed. Please check the model file and ensure it's properly placed in the text_encoders or clip folder." print(error_msg) return (error_msg,)