diff --git a/.gitignore b/.gitignore
new file mode 100644
index 0000000..8496516
--- /dev/null
+++ b/.gitignore
@@ -0,0 +1,29 @@
+# Test files
+test_*.py
+*_test.py
+
+# Test markdown files (except README.md)
+*.md
+!README.md
+
+# Config files
+openaimodel.json
+
+# Python cache
+__pycache__/
+*.pyc
+*.pyo
+*.pyd
+.Python
+
+# IDE
+.vscode/
+.idea/
+*.swp
+*.swo
+
+# Prompt directory
+Prompt/
+
+# Claude AI directory
+.claude/
diff --git a/nodes.py b/nodes.py
index 1eb9405..b4dc346 100644
--- a/nodes.py
+++ b/nodes.py
@@ -23,8 +23,6 @@ from comfy.cli_args import args
from typing import List, Dict, Any, Tuple
from random import Random
from datetime import datetime
-from .openrouter_llm import OpenRouterLLM
-from .openai_helper import OpenAIHelper
from .qwen_inference import QwenGPUInference
class AudioListGenerator:
@@ -938,8 +936,6 @@ NODE_CLASS_MAPPINGS = {
"LoadVideoPath": LoadVideoPath,
"SaveVideoPath": SaveVideoPath,
"FrameMatch": FrameMatch,
- "OpenRouterLLM": OpenRouterLLM,
- "OpenAIHelper": OpenAIHelper,
"QwenGPUInference": QwenGPUInference,
}
@@ -953,8 +949,6 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"LoadVideoPath": "LoadVideoPath",
"SaveVideoPath": "SaveVideoPath",
"FrameMatch": "FrameMatch",
- "OpenRouterLLM": "OpenRouter LLM",
- "OpenAIHelper": "OpenAI Helper",
"QwenGPUInference": "Qwen GPU Inference",
}
diff --git a/openai_helper.py b/openai_helper.py
new file mode 100644
index 0000000..4b7bf5f
--- /dev/null
+++ b/openai_helper.py
@@ -0,0 +1,423 @@
+import requests
+import json
+import base64
+import os
+import io
+import torch
+from PIL import Image
+import numpy as np
+import tempfile
+
+# 檢查 torchaudio 是否可用
+TORCHAUDIO_AVAILABLE = False
+try:
+ import torchaudio
+ TORCHAUDIO_AVAILABLE = True
+except ImportError:
+ print("⚠️ torchaudio 未安裝,音訊功能將無法使用")
+
+class OpenAIHelper:
+ """
+ OpenAI Helper節點,用於呼叫OpenAI相容的API
+ 支援圖片、音訊輸入,配置管理,模型列表獲取
+ """
+
+ @classmethod
+ def _load_config(cls):
+ """載入配置從openaimodel.json文件"""
+ config_file = os.path.join(os.path.dirname(__file__), "openaimodel.json")
+ try:
+ with open(config_file, 'r', encoding='utf-8') as f:
+ return json.load(f)
+ except Exception as e:
+ print(f"⚠️ 載入openaimodel.json失敗: {e}")
+ return {
+ "endpoint": "",
+ "api_key": "",
+ "model_name": ""
+ }
+
+ @classmethod
+ def _save_config(cls, endpoint, api_key, model_name):
+ """保存配置到openaimodel.json文件"""
+ config_file = os.path.join(os.path.dirname(__file__), "openaimodel.json")
+ try:
+ config = {
+ "endpoint": endpoint,
+ "api_key": api_key,
+ "model_name": model_name
+ }
+ with open(config_file, 'w', encoding='utf-8') as f:
+ json.dump(config, f, ensure_ascii=False, indent=2)
+ except Exception as e:
+ print(f"❌ 保存openaimodel.json失敗: {e}")
+
+ @classmethod
+ def INPUT_TYPES(cls):
+ # 載入保存的配置
+ config = cls._load_config()
+
+ return {
+ "required": {
+ "endpoint": ("STRING", {
+ "multiline": False,
+ "default": config.get("endpoint", ""),
+ "placeholder": "輸入OpenAI API端點,例如: https://api.openai.com/v1/chat/completions"
+ }),
+ "api_key": ("STRING", {
+ "multiline": False,
+ "default": config.get("api_key", ""),
+ "placeholder": "輸入您的API金鑰"
+ }),
+ "model_name": ("STRING", {
+ "multiline": False,
+ "default": config.get("model_name", ""),
+ "placeholder": "輸入模型名稱,例如: gpt-4o"
+ }),
+ "user_prompt": ("STRING", {
+ "multiline": True,
+ "default": "請分析提供的內容。"
+ }),
+ "max_tokens": ("INT", {
+ "default": 2000,
+ "min": 1,
+ "max": 128000,
+ "step": 1
+ }),
+ },
+ "optional": {
+ "system_prompt": ("STRING", {
+ "multiline": True,
+ "default": "請以繁體中文輸出使用者內容,不須包括引導或後綴,如「這就是你要的結果」、「以下是你要的結果」、「你要不要我幫你」、「你說的對」等等,只需要輸出使用者要的結論raw_text。請勿使用Markdown語法(如**粗體**),直接輸出純文字即可。"
+ }),
+ "image1": ("IMAGE",),
+ "image2": ("IMAGE",),
+ "image3": ("IMAGE",),
+ "audio": ("AUDIO",),
+ "file_path": ("STRING", {"multiline": False, "default": ""}),
+ }
+ }
+
+ RETURN_TYPES = ("STRING", "STRING")
+ RETURN_NAMES = ("text", "model_name_list")
+ FUNCTION = "process_openai"
+ CATEGORY = "ListHelper"
+
+ def __init__(self):
+ pass
+
+ def _process_audio(self, audio):
+ """處理音訊並轉換為base64編碼"""
+ if not TORCHAUDIO_AVAILABLE:
+ print("❌ torchaudio未安裝,無法處理音訊")
+ return None
+
+ try:
+ temp_file = None
+
+ # 檢查不同的音訊輸入格式
+ if isinstance(audio, dict):
+ if "path" in audio:
+ # 直接路徑格式
+ audio_path = audio["path"]
+ print(f"處理來自路徑的音訊: {audio_path}")
+
+ if not os.path.exists(audio_path):
+ print(f"❌ 音訊文件不存在: {audio_path}")
+ return None
+
+ # 讀取音訊文件並轉換為base64
+ with open(audio_path, 'rb') as f:
+ audio_bytes = f.read()
+ return base64.b64encode(audio_bytes).decode('utf-8')
+
+ elif "waveform" in audio and "sample_rate" in audio:
+ # ComfyUI音訊節點格式
+ print(f"處理來自waveform tensor的音訊")
+ waveform = audio["waveform"]
+ sample_rate = audio["sample_rate"]
+
+ # 創建臨時WAV文件
+ temp_file = tempfile.NamedTemporaryFile(delete=False, suffix='.wav')
+ temp_path = temp_file.name
+ temp_file.close()
+
+ # 確保waveform格式正確 [channels, samples]
+ if waveform.dim() == 3:
+ waveform = waveform.squeeze(0) # 移除批次維度
+
+ # 保存為WAV文件
+ torchaudio.save(temp_path, waveform, sample_rate)
+
+ # 讀取並轉換為base64
+ with open(temp_path, 'rb') as f:
+ audio_bytes = f.read()
+ audio_b64 = base64.b64encode(audio_bytes).decode('utf-8')
+
+ # 清理臨時文件
+ os.unlink(temp_path)
+ return audio_b64
+
+ else:
+ # 未知字典格式
+ print(f"❌ 未知的音訊字典格式: {list(audio.keys())}")
+ return None
+
+ elif isinstance(audio, str) and os.path.exists(audio):
+ # 直接文件路徑
+ print(f"處理來自直接路徑的音訊: {audio}")
+ with open(audio, 'rb') as f:
+ audio_bytes = f.read()
+ return base64.b64encode(audio_bytes).decode('utf-8')
+
+ else:
+ # 嘗試作為tensor處理
+ print(f"嘗試將音訊作為tensor處理")
+ if hasattr(audio, 'shape'):
+ # 創建臨時WAV文件
+ temp_file = tempfile.NamedTemporaryFile(delete=False, suffix='.wav')
+ temp_path = temp_file.name
+ temp_file.close()
+
+ # 假設sample_rate為44100(可以調整)
+ sample_rate = 44100
+ torchaudio.save(temp_path, audio, sample_rate)
+
+ # 讀取並轉換為base64
+ with open(temp_path, 'rb') as f:
+ audio_bytes = f.read()
+ audio_b64 = base64.b64encode(audio_bytes).decode('utf-8')
+
+ # 清理臨時文件
+ os.unlink(temp_path)
+ return audio_b64
+ else:
+ print(f"❌ 無法識別的音訊格式")
+ return None
+
+ except Exception as e:
+ print(f"❌ 處理音訊時出錯: {e}")
+ if temp_file and os.path.exists(temp_file.name):
+ os.unlink(temp_file.name)
+ return None
+
+ def _tensor_to_base64(self, tensor):
+ """將ComfyUI圖像tensor轉換為base64編碼"""
+ # tensor shape: [B, H, W, C] (0-1 range)
+ if tensor.dim() == 4:
+ tensor = tensor.squeeze(0) # 移除批次維度
+
+ # 轉換為numpy並調整範圍到0-255
+ numpy_image = (tensor.cpu().numpy() * 255).astype(np.uint8)
+
+ # 轉換為PIL圖像
+ pil_image = Image.fromarray(numpy_image)
+
+ # 轉換為base64
+ buffer = io.BytesIO()
+ pil_image.save(buffer, format='PNG')
+ image_base64 = base64.b64encode(buffer.getvalue()).decode('utf-8')
+
+ return f"data:image/png;base64,{image_base64}"
+
+ def _get_model_list(self, endpoint, api_key):
+ """獲取可用模型列表"""
+ try:
+ # 將 chat/completions 端點改為 models 端點
+ base_url = endpoint.rsplit('/chat/completions', 1)[0]
+ models_endpoint = f"{base_url}/models"
+
+ headers = {
+ "Authorization": f"Bearer {api_key}",
+ "Content-Type": "application/json"
+ }
+
+ response = requests.get(
+ models_endpoint,
+ headers=headers,
+ timeout=10
+ )
+
+ if response.status_code == 200:
+ data = response.json()
+ if 'data' in data:
+ # 提取模型ID
+ models = [model.get('id', '') for model in data['data']]
+ return ', '.join(models)
+ else:
+ return "無法獲取模型列表"
+ else:
+ return f"獲取模型列表失敗: HTTP {response.status_code}"
+
+ except Exception as e:
+ print(f"❌ 獲取模型列表異常: {e}")
+ return f"獲取模型列表異常: {str(e)}"
+
+ def _call_openai_api(self, endpoint, api_key, model, messages, max_tokens, audio_b64=None):
+ """呼叫OpenAI API"""
+ headers = {
+ "Authorization": f"Bearer {api_key}",
+ "Content-Type": "application/json"
+ }
+
+ data = {
+ "model": model,
+ "messages": messages,
+ "max_tokens": max_tokens
+ }
+
+ try:
+ response = requests.post(
+ endpoint,
+ headers=headers,
+ json=data,
+ timeout=60
+ )
+
+ if response.status_code == 200:
+ result = response.json()
+ return result
+ else:
+ try:
+ error_data = response.json()
+ return {"error": error_data.get("error", {"message": f"HTTP {response.status_code}"})}
+ except:
+ return {"error": {"message": f"HTTP {response.status_code}: {response.text[:200]}"}}
+
+ except Exception as e:
+ print(f"❌ API呼叫異常: {e}")
+ return {"error": {"message": str(e)}}
+
+ def process_openai(self, endpoint, api_key, model_name, user_prompt, max_tokens,
+ system_prompt=None, image1=None, image2=None, image3=None,
+ audio=None, file_path=None):
+ """處理OpenAI請求"""
+
+ # 驗證必填參數
+ if not endpoint or not endpoint.strip():
+ return ("❌ 錯誤: 請提供API端點", "")
+
+ if not api_key or not api_key.strip():
+ return ("❌ 錯誤: 請提供API金鑰", "")
+
+ if not model_name or not model_name.strip():
+ return ("❌ 錯誤: 請提供模型名稱", "")
+
+ # 保存配置
+ self._save_config(endpoint.strip(), api_key.strip(), model_name.strip())
+
+ # 獲取模型列表
+ model_list = self._get_model_list(endpoint.strip(), api_key.strip())
+
+ # 準備消息內容
+ messages = []
+
+ # 添加系統提示(如果有的話)
+ if system_prompt and system_prompt.strip():
+ messages.append({"role": "system", "content": system_prompt.strip()})
+
+ # 構建用戶消息
+ user_content = []
+
+ # 添加文字內容
+ if user_prompt and user_prompt.strip():
+ user_content.append({
+ "type": "text",
+ "text": user_prompt.strip()
+ })
+
+ # 檢查是否有圖像輸入
+ images_to_process = []
+ for img_input in [image1, image2, image3]:
+ if img_input is not None:
+ images_to_process.append(img_input)
+
+ # 添加圖像到用戶消息
+ for image_tensor in images_to_process:
+ try:
+ base64_image = self._tensor_to_base64(image_tensor)
+ user_content.append({
+ "type": "image_url",
+ "image_url": {
+ "url": base64_image
+ }
+ })
+ except Exception as e:
+ print(f"❌ 處理圖像失敗: {e}")
+
+ # 處理音訊輸入
+ audio_b64 = None
+ if audio is not None:
+ if not TORCHAUDIO_AVAILABLE:
+ return ("❌ 錯誤: torchaudio未安裝,無法處理音訊", model_list)
+
+ print(f"處理音訊輸入")
+ try:
+ audio_b64 = self._process_audio(audio)
+ if audio_b64:
+ # 添加音訊到用戶消息(使用inline_data格式)
+ user_content.append({
+ "type": "input_audio",
+ "input_audio": {
+ "data": audio_b64,
+ "format": "wav"
+ }
+ })
+ print(f"✓ 音訊處理成功")
+ else:
+ return ("❌ 錯誤: 音訊處理失敗", model_list)
+ except Exception as e:
+ print(f"❌ 處理音訊時出錯: {str(e)}")
+ return (f"❌ 錯誤: 處理音訊時出錯: {str(e)}", model_list)
+
+ # 添加用戶消息
+ if len(user_content) > 1:
+ # 多媒體內容(包含圖像或音訊)
+ messages.append({
+ "role": "user",
+ "content": user_content
+ })
+ elif user_content:
+ # 純文字內容
+ messages.append({
+ "role": "user",
+ "content": user_content[0]["text"]
+ })
+ else:
+ # 空內容
+ messages.append({
+ "role": "user",
+ "content": "請分析提供的內容。"
+ })
+
+ # 呼叫API
+ response = self._call_openai_api(
+ endpoint.strip(),
+ api_key.strip(),
+ model_name.strip(),
+ messages,
+ max_tokens,
+ audio_b64
+ )
+
+ # 處理響應
+ if not response:
+ return ("❌ API呼叫失敗", model_list)
+
+ # 檢查錯誤
+ if 'error' in response:
+ error_info = response['error']
+ if isinstance(error_info, dict):
+ error_msg = f"❌ API錯誤: {error_info.get('message', str(error_info))}"
+ else:
+ error_msg = f"❌ API錯誤: {str(error_info)}"
+ return (error_msg, model_list)
+
+ if 'choices' not in response or not response['choices']:
+ return ("❌ API回應格式異常", model_list)
+
+ # 獲取回應訊息
+ message = response['choices'][0]['message']
+ response_content = message.get('content', '')
+
+ return (response_content, model_list)
diff --git a/qwen_inference.py b/qwen_inference.py
index 569929c..3cc80c5 100644
--- a/qwen_inference.py
+++ b/qwen_inference.py
@@ -9,9 +9,9 @@ from typing import Optional, Tuple, Dict
class QwenGPUInference:
"""
- Qwen3-4B GPU 推理節點(優化載入速度版本 v3 - 支援記憶體管理)
- 自動下載所需配置檔案並使用 GPU 進行推理
- 包含 GPU 記憶體檢查與清理功能,避免與 ComfyUI 的 CLIP 模型衝突
+ Qwen3-4B GPU Inference Node with intelligent memory management
+ Auto-downloads required config files and performs GPU inference
+ Includes GPU memory checking and cleanup to avoid conflicts with ComfyUI's CLIP models
"""
def __init__(self):
@@ -22,7 +22,7 @@ class QwenGPUInference:
@classmethod
def _get_safetensors_files(cls):
- """從 text_encoders 資料夾中獲取所有 safetensors 檔案"""
+ """Get all safetensors files from text_encoders folder"""
safetensors_files = []
try:
@@ -42,38 +42,60 @@ class QwenGPUInference:
return sorted(safetensors_files)
+ @classmethod
+ def _get_prompt_templates(cls):
+ """Get all .md template files from Prompt folder"""
+ current_dir = os.path.dirname(os.path.abspath(__file__))
+ prompt_dir = os.path.join(current_dir, "Prompt")
+
+ templates = []
+
+ if os.path.exists(prompt_dir):
+ for file in os.listdir(prompt_dir):
+ if file.lower().endswith('.md'):
+ templates.append(file)
+
+ if not templates:
+ return ["No Template"]
+
+ return sorted(templates)
+
+ @classmethod
+ def _load_template_content(cls, template_name):
+ """Load template content"""
+ if template_name == "No Template" or template_name == "Custom":
+ return ""
+
+ current_dir = os.path.dirname(os.path.abspath(__file__))
+ template_path = os.path.join(current_dir, "Prompt", template_name)
+
+ if os.path.exists(template_path):
+ try:
+ with open(template_path, 'r', encoding='utf-8') as f:
+ return f.read()
+ except:
+ return ""
+
+ return ""
+
@classmethod
def INPUT_TYPES(cls):
+ templates = cls._get_prompt_templates()
+ # Add "Custom" option at the beginning of the list
+ template_options = ["Custom"] + templates
+
return {
"required": {
"user_prompt": ("STRING", {
"multiline": True,
- "default": "一個女孩在咖啡廳"
+ "default": "A girl in a coffee shop"
+ }),
+ "prompt_template": (template_options, {
+ "default": template_options[0] if template_options else "Custom"
}),
"system_prompt": ("STRING", {
"multiline": True,
- "default": """你是一位專業的攝影提示詞優化專家。你的任務是將簡單的場景描述轉換為詳細、專業的攝影提示詞。
-
-請根據用戶輸入的簡單描述,生成包含以下元素的完整提示詞:
-
-1. **主體描述**:詳細描述主要拍攝對象(人物、物體、場景)
-2. **環境細節**:周圍環境、背景元素、場景氛圍
-3. **光影效果**:光線類型(自然光/人造光)、光線方向、光影對比、色溫
-4. **相機設定**:視角、景深、焦距效果
-5. **構圖元素**:畫面佈局、前景/中景/背景關係
-6. **色彩氛圍**:主色調、色彩搭配、飽和度
-7. **質感細節**:材質、紋理、細節表現
-8. **情緒氛圍**:整體氛圍、情感表達
-
-輸出格式:
-- 使用英文輸出專業攝影術語
-- 用逗號分隔各個元素
-- 確保描述具體、可視覺化
-- 長度控制在 150-300 個英文單詞
-
-範例:
-輸入:一個女孩在咖啡廳
-輸出:A young woman sitting by the window in a cozy coffee shop, warm afternoon sunlight streaming through large glass windows creating soft shadows, wearing casual outfit, holding a cup of coffee, wooden table with laptop and notebook, blurred background with other customers, shallow depth of field, bokeh effect, warm color temperature, golden hour lighting, natural skin tones, professional photography, shot with 50mm lens, f/1.8 aperture, Instagram aesthetic, lifestyle photography, candid moment, peaceful atmosphere"""
+ "default": ""
}),
"max_new_tokens": ("INT", {
"default": 2048,
@@ -91,7 +113,7 @@ class QwenGPUInference:
"optional": {
"do_sample": ("BOOLEAN", {
"default": True,
- "tooltip": "是否使用採樣"
+ "tooltip": "Enable sampling"
}),
"top_p": ("FLOAT", {
"default": 0.9,
@@ -114,42 +136,42 @@ class QwenGPUInference:
CATEGORY = "ListHelper"
def _find_qwen_model(self) -> Optional[str]:
- """自動尋找 qwen_3_4b.safetensors 模型"""
+ """Auto-find qwen_3_4b.safetensors model"""
safetensors_files = self._get_safetensors_files()
- # 優先尋找 qwen_3_4b.safetensors
+ # Prioritize finding qwen_3_4b.safetensors
for path in safetensors_files:
if path != "No safetensors files found":
basename = os.path.basename(path).lower()
if "qwen" in basename and "3" in basename and "4b" in basename:
return path
- # 如果找不到特定模型,返回第一個 safetensors 檔案
+ # If specific model not found, return first safetensors file
if safetensors_files and safetensors_files[0] != "No safetensors files found":
return safetensors_files[0]
return None
def _remove_thinking_tags(self, text: str) -> str:
- """移除 ... 標籤及其內容"""
- # 使用正則表達式移除所有 ... 區塊
+ """Remove ... tags and their content"""
+ # Use regex to remove all ... blocks
cleaned_text = re.sub(r'.*?', '', text, flags=re.DOTALL)
- # 移除多餘的空白行
+ # Remove extra blank lines
cleaned_text = re.sub(r'\n\s*\n', '\n', cleaned_text)
return cleaned_text.strip()
def _check_gpu_memory(self, required_gb: float = 8.0) -> Tuple[bool, str]:
"""
- 檢查 GPU 記憶體是否足夠
+ Check if GPU memory is sufficient
Args:
- required_gb: 需要的 GPU 記憶體大小(GB)
+ required_gb: Required GPU memory size (GB)
Returns:
- (是否足夠, 詳細訊息)
+ (is_sufficient, detailed_message)
"""
if not torch.cuda.is_available():
- return True, "使用 CPU 模式,無需檢查 GPU 記憶體"
+ return True, "GPU Memory: Using CPU mode"
try:
total_memory = torch.cuda.get_device_properties(0).total_memory / 1024**3
@@ -157,67 +179,53 @@ class QwenGPUInference:
reserved_memory = torch.cuda.memory_reserved(0) / 1024**3
free_memory = total_memory - reserved_memory
- info = f"""
-GPU 記憶體狀態:
- 總記憶體: {total_memory:.2f} GB
- 已分配: {allocated_memory:.2f} GB
- 已保留: {reserved_memory:.2f} GB
- 可用: {free_memory:.2f} GB
- 需要: {required_gb:.2f} GB
-"""
+ info = f"GPU Memory: Total {total_memory:.2f}GB | Free {free_memory:.2f}GB | Required {required_gb:.2f}GB"
if free_memory < required_gb:
- return False, info + f"\n⚠️ 記憶體不足!缺少 {required_gb - free_memory:.2f} GB"
+ return False, info + f" | Insufficient: need {required_gb - free_memory:.2f}GB more"
- return True, info + "\n✓ 記憶體充足"
+ return True, info + " | Sufficient"
except Exception as e:
- return True, f"無法檢查 GPU 記憶體: {e}"
+ return True, f"GPU Memory: Cannot check - {e}"
def _free_gpu_memory(self) -> None:
"""
- 釋放 GPU 記憶體
- 清理 PyTorch 快取和執行垃圾回收
+ Free GPU memory
+ Clear PyTorch cache and run garbage collection
"""
try:
if torch.cuda.is_available():
- # 記錄清理前的記憶體
- before_allocated = torch.cuda.memory_allocated(0) / 1024**3
+ # Record memory before cleanup
before_reserved = torch.cuda.memory_reserved(0) / 1024**3
- # 清理 CUDA 快取
+ # Clear CUDA cache
torch.cuda.empty_cache()
torch.cuda.synchronize()
- # 強制垃圾回收
+ # Force garbage collection
gc.collect()
- # 再次清理
+ # Clear again
torch.cuda.empty_cache()
- # 記錄清理後的記憶體
- after_allocated = torch.cuda.memory_allocated(0) / 1024**3
+ # Record memory after cleanup
after_reserved = torch.cuda.memory_reserved(0) / 1024**3
- freed_allocated = before_allocated - after_allocated
freed_reserved = before_reserved - after_reserved
- print(f"\n✓ GPU 記憶體已清理:")
- print(f" 釋放已分配記憶體: {freed_allocated:.2f} GB")
- print(f" 釋放已保留記憶體: {freed_reserved:.2f} GB")
- print(f" 當前已分配: {after_allocated:.2f} GB")
- print(f" 當前已保留: {after_reserved:.2f} GB")
+ print(f"GPU Memory: Freed {freed_reserved:.2f}GB | Current reserved {after_reserved:.2f}GB")
else:
gc.collect()
- print("✓ 執行垃圾回收(CPU 模式)")
+ print("Memory cleanup: CPU mode")
except Exception as e:
- print(f"⚠️ 清理記憶體時發生錯誤: {e}")
- # 即使發生錯誤,仍嘗試垃圾回收
+ print(f"Memory cleanup error: {e}")
+ # Attempt garbage collection even if error occurs
gc.collect()
def _download_config_files(self, model_path: str, repo_id: str) -> bool:
- """自動下載 HuggingFace 配置檔案"""
+ """Auto-download HuggingFace config files"""
try:
model_dir = os.path.dirname(model_path)
model_basename = os.path.splitext(os.path.basename(model_path))[0]
@@ -235,11 +243,11 @@ GPU 記憶體狀態:
all_exist = all(os.path.exists(os.path.join(self.config_dir, f)) for f in config_files)
if all_exist:
- print(f"✓ 配置檔案已存在: {self.config_dir}")
+ print(f"Config files exist: {self.config_dir}")
return True
os.makedirs(self.config_dir, exist_ok=True)
- print(f"下載配置檔案到: {self.config_dir}")
+ print(f"Downloading config files to: {self.config_dir}")
base_url = f"https://huggingface.co/{repo_id}/resolve/main/"
@@ -247,11 +255,11 @@ GPU 記憶體狀態:
filepath = os.path.join(self.config_dir, filename)
if os.path.exists(filepath):
- print(f" ✓ {filename} 已存在")
+ print(f" {filename} exists")
continue
url = base_url + filename
- print(f" 下載 {filename}...")
+ print(f" Downloading {filename}...")
try:
response = requests.get(url, timeout=30)
@@ -259,129 +267,121 @@ GPU 記憶體狀態:
with open(filepath, 'wb') as f:
f.write(response.content)
- print(f" ✓ {filename} 下載完成")
+ print(f" {filename} downloaded")
except Exception as e:
- print(f" ✗ {filename} 下載失敗: {str(e)}")
+ print(f" Failed to download {filename}: {str(e)}")
continue
return True
except Exception as e:
- print(f"❌ 下載配置檔案失敗: {e}")
+ print(f"Failed to download config files: {e}")
import traceback
traceback.print_exc()
return False
def _load_model(self, model_path: str, repo_id: str) -> bool:
- """載入模型和 tokenizer(優化版 v3 - 包含記憶體管理)"""
+ """Load model and tokenizer (optimized v3 - with memory management)"""
try:
if self.model is not None and self.current_model_path == model_path:
- print(f"✓ 模型已載入: {os.path.basename(model_path)}")
+ print(f"Model already loaded: {os.path.basename(model_path)}")
return True
try:
from transformers import AutoModelForCausalLM, AutoTokenizer, AutoConfig
from safetensors.torch import load_file
except ImportError as e:
- print(f"❌ 缺少必要的套件: {e}")
- print("請執行: pip install transformers safetensors")
+ print(f"Missing required packages: {e}")
+ print("Please run: pip install transformers safetensors")
return False
if not self._download_config_files(model_path, repo_id):
return False
- print(f"\n載入模型: {os.path.basename(model_path)}")
- print("=" * 80)
+ print(f"\nLoading model: {os.path.basename(model_path)}")
- # 步驟 1: 檢查 GPU 記憶體
- print("\n步驟 1/3: 檢查 GPU 記憶體...")
+ # Step 1: Check GPU memory
is_enough, memory_info = self._check_gpu_memory(required_gb=7.5)
print(memory_info)
- # 步驟 2: 如果記憶體不足,嘗試清理
+ # Step 2: If insufficient memory, try to free up
if not is_enough:
- print("\n步驟 2/3: 記憶體不足,執行清理...")
+ print("Memory insufficient, cleaning up...")
self._free_gpu_memory()
- # 再次檢查
+ # Check again
is_enough, memory_info = self._check_gpu_memory(required_gb=7.5)
- print("\n清理後的記憶體狀態:")
print(memory_info)
if not is_enough:
- print("\n⚠️ GPU 記憶體仍然不足,將使用 CPU Offload 策略")
- print(" - 部分模型層會放在 CPU,推理速度會較慢")
- else:
- print("\n步驟 2/3: 記憶體充足,跳過清理")
+ print("Still insufficient, will use CPU Offload strategy (slower inference)")
- print("\n步驟 3/3: 載入模型...")
overall_start = time.time()
device = "cuda:0" if torch.cuda.is_available() else "cpu"
dtype = torch.float16 if torch.cuda.is_available() else torch.float32
- # 1. 載入 tokenizer
+ # 1. Load tokenizer
self.tokenizer = AutoTokenizer.from_pretrained(
self.config_dir,
trust_remote_code=True
)
- # 2. 載入配置
+ # 2. Load config
config = AutoConfig.from_pretrained(self.config_dir, trust_remote_code=True)
- # 3. 根據記憶體情況選擇載入策略
+ # 3. Choose loading strategy based on memory
import shutil
temp_model_dir = os.path.join(self.config_dir, "temp_model")
os.makedirs(temp_model_dir, exist_ok=True)
- # 決定載入策略
+ # 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
if free_memory >= 7.5:
- # 充足記憶體:完全載入到 GPU
- print(f"⚡ 策略 1: 完全 GPU 載入(可用記憶體: {free_memory:.2f} GB)")
+ # Sufficient memory: Full GPU loading
+ print(f"Strategy: Full GPU loading (Free: {free_memory:.2f}GB)")
max_memory_config = None
offload_folder = None
else:
- # 記憶體不足:使用 CPU Offload
- available_gpu = max(3.0, free_memory - 1.0) # 至少保留 1GB 給其他操作
- print(f"⚡ 策略 2: CPU Offload(可用 GPU: {free_memory:.2f} GB,分配: {available_gpu:.2f} GB)")
- print(f" - 部分模型層將放在 CPU,推理速度會較慢")
+ # Insufficient memory: Use CPU Offload
+ available_gpu = max(3.0, free_memory - 1.0) # Reserve at least 1GB
+ print(f"Strategy: CPU Offload (Free: {free_memory:.2f}GB, Allocate: {available_gpu:.2f}GB)")
max_memory_config = {
0: f"{available_gpu:.1f}GB",
"cpu": "16GB"
}
- # 創建 offload 資料夾
+ # Create offload folder
offload_folder = os.path.join(self.config_dir, "offload")
os.makedirs(offload_folder, exist_ok=True)
else:
- print("⚡ 策略 3: CPU 模式")
+ print("Strategy: CPU mode")
max_memory_config = None
offload_folder = None
load_start = time.time()
try:
- # 複製配置檔案
+ # Copy config files
for file in ["config.json", "generation_config.json"]:
src = os.path.join(self.config_dir, file)
if os.path.exists(src):
shutil.copy(src, temp_model_dir)
- # 創建符號連結或複製 safetensors
+ # Create link or copy safetensors
safetensors_target = os.path.join(temp_model_dir, "model.safetensors")
if os.path.exists(safetensors_target):
os.remove(safetensors_target)
- # Windows 使用硬連結而不是符號連結
+ # Windows uses hard links instead of symbolic links
try:
os.link(model_path, safetensors_target)
except:
shutil.copy(model_path, safetensors_target)
- # 使用適當的載入策略
+ # Use appropriate loading strategy
load_kwargs = {
"pretrained_model_name_or_path": temp_model_dir,
"trust_remote_code": True,
@@ -398,64 +398,29 @@ GPU 記憶體狀態:
self.model = AutoModelForCausalLM.from_pretrained(**load_kwargs)
- print(f" ✓ 模型載入完成(耗時: {time.time() - load_start:.2f} 秒)")
+ print(f"Model loaded (Time: {time.time() - load_start:.2f}s)")
finally:
- # 清理臨時目錄
+ # Clean up temp directory
try:
if os.path.exists(temp_model_dir):
shutil.rmtree(temp_model_dir)
except:
pass
- missing_keys = []
- unexpected_keys = []
-
- if missing_keys:
- print(f" 警告: 缺少的鍵值: {len(missing_keys)} 個")
- if unexpected_keys:
- print(f" 警告: 未預期的鍵值: {len(unexpected_keys)} 個")
-
self.current_model_path = model_path
total_time = time.time() - overall_start
- print(f"\n驗證模型狀態...")
if torch.cuda.is_available():
current_allocated = torch.cuda.memory_allocated(0) / 1024**3
current_reserved = torch.cuda.memory_reserved(0) / 1024**3
- print(f" ✓ 模型已載入")
- print(f" - GPU 記憶體已分配: {current_allocated:.2f} GB")
- print(f" - GPU 記憶體已保留: {current_reserved:.2f} GB")
+ print(f"GPU Memory: Allocated {current_allocated:.2f}GB | Reserved {current_reserved:.2f}GB")
- # 檢查模型設備分佈
- device_map = {}
- for name, param in self.model.named_parameters():
- device_str = str(param.device)
- device_map[device_str] = device_map.get(device_str, 0) + 1
-
- print(f" - 模型設備分佈:")
- for device_name, count in device_map.items():
- print(f" * {device_name}: {count} 個參數")
- else:
- print(f" ✓ 模型在 CPU 上")
-
- print("\n" + "=" * 80)
- print(f"✓ 模型載入成功(總耗時: {total_time:.2f} 秒)")
- print("=" * 80)
-
- # 給出優化建議
- if total_time > 60:
- print(f"\n💡 載入優化建議:")
- print(f" - 當前載入時間: {total_time:.1f} 秒")
- print(f" - 主要瓶頸: 移動模型到 GPU")
- print(f" - 這是正常的,無法進一步優化(硬體限制)")
- print(f" - 模型會保留在記憶體中,下次使用會即時載入")
-
- print()
+ print(f"Model loaded successfully (Total time: {total_time:.2f}s)")
return True
except Exception as e:
- print(f"❌ 載入模型失敗: {e}")
+ print(f"Failed to load model: {e}")
import traceback
traceback.print_exc()
self.model = None
@@ -466,6 +431,7 @@ GPU 記憶體狀態:
def inference(
self,
user_prompt: str,
+ prompt_template: str,
system_prompt: str,
max_new_tokens: int,
temperature: float,
@@ -473,22 +439,37 @@ GPU 記憶體狀態:
top_p: float = 0.9,
top_k: int = 50,
) -> Tuple[str]:
- """執行推理"""
+ """Execute inference"""
- # 自動尋找 Qwen 模型
+ # Auto-find Qwen model
model_path = self._find_qwen_model()
if model_path is None:
- return ("❌ 錯誤: 未找到 Qwen 模型檔案\n請將 qwen_3_4b.safetensors 檔案放在 ComfyUI 的 models/text_encoders 資料夾中",)
+ error_msg = "Error: Qwen model file not found.\nPlease place the correct model file (e.g., qwen_3_4b.safetensors) in ComfyUI's models/text_encoders folder."
+ print(error_msg)
+ return (error_msg,)
if not os.path.exists(model_path):
- return (f"❌ 錯誤: 模型檔案不存在: {model_path}",)
+ error_msg = f"Error: Model file does not exist: {model_path}\nPlease place the correct model file in the text_encoders folder."
+ print(error_msg)
+ return (error_msg,)
- # 使用固定的 repo_id
+ # Use fixed repo_id
repo_id = "Qwen/Qwen3-4B"
if not self._load_model(model_path, repo_id):
- return ("❌ 錯誤: 模型載入失敗",)
+ error_msg = "Error: Model loading failed. Please check the model file and ensure it's properly placed in the text_encoders folder."
+ print(error_msg)
+ return (error_msg,)
try:
+ # Load and apply template
+ template_content = ""
+ if prompt_template != "Custom":
+ template_content = self._load_template_content(prompt_template)
+
+ # If template content exists, replace system_prompt with template
+ if template_content:
+ system_prompt = template_content
+
messages = []
if system_prompt and system_prompt.strip():
messages.append({"role": "system", "content": system_prompt})
@@ -506,7 +487,7 @@ GPU 記憶體狀態:
inputs = {k: v.to(device) for k, v in inputs.items()}
inference_start = time.time()
- print(f"開始推理...")
+ print(f"Inference starting...")
with torch.no_grad():
outputs = self.model.generate(
@@ -522,14 +503,14 @@ GPU 記憶體狀態:
response = self.tokenizer.decode(outputs[0], skip_special_tokens=True)
- # 移除 assistant 標記
+ # Remove assistant markers
if "assistant" in response:
for separator in ["<|im_start|>assistant\n", "assistant\n", "Assistant:", "assistant:"]:
if separator in response:
response = response.split(separator)[-1].strip()
break
- # 移除 標籤
+ # Remove tags
response = self._remove_thinking_tags(response)
inference_time = time.time() - inference_start
@@ -537,18 +518,15 @@ GPU 記憶體狀態:
tokens_generated = len(outputs[0]) - len(inputs['input_ids'][0])
tokens_per_sec = tokens_generated / inference_time if inference_time > 0 else 0
- print(f"✓ 推理完成(耗時: {inference_time:.2f} 秒)")
- print(f" 生成 tokens: {tokens_generated}")
- print(f" 速度: {tokens_per_sec:.1f} tokens/秒")
+ print(f"Inference completed (Time: {inference_time:.2f}s | Tokens: {tokens_generated} | Speed: {tokens_per_sec:.1f} tokens/s)")
if torch.cuda.is_available():
- print(f" GPU 記憶體使用: {torch.cuda.memory_allocated(0) / 1024**3:.2f} GB")
- print(f" GPU 記憶體峰值: {torch.cuda.max_memory_allocated(0) / 1024**3:.2f} GB")
+ print(f"GPU Memory: Used {torch.cuda.memory_allocated(0) / 1024**3:.2f}GB | Peak {torch.cuda.max_memory_allocated(0) / 1024**3:.2f}GB")
return (response,)
except Exception as e:
import traceback
- error_msg = f"❌ 推理失敗: {str(e)}\n{traceback.format_exc()}"
+ error_msg = f"Inference failed: {str(e)}\n{traceback.format_exc()}"
print(error_msg)
return (error_msg,)
diff --git a/requirements.txt b/requirements.txt
index 822e14a..b4a9a3a 100644
--- a/requirements.txt
+++ b/requirements.txt
@@ -1 +1,5 @@
-regex
\ No newline at end of file
+regex
+accelerate
+transformers
+safetensors
+requests
\ No newline at end of file