@@ -0,0 +1,16 @@
|
||||
# Repository Guidelines
|
||||
|
||||
## Project Structure & Module Organization
|
||||
Core logic lives in `nodes.py`, which defines the Qwen multimodal and text nodes plus helper utilities for CUDA memory and Hugging Face downloads. Packaging metadata is in `pyproject.toml`, dependencies in `requirements.txt`, and reusable launch configs under `config/`. Reference workflows (`workflow/*.json`) demonstrate expected node wiring and ship with screenshots for quick visual checks. Tests mirror the runtime entry points inside `tests/`, while ComfyUI caches downloaded checkpoints in `ComfyUI/models/LLM/`; avoid committing anything below that directory.
|
||||
|
||||
## Build, Test, and Development Commands
|
||||
Install dependencies inside the ComfyUI virtualenv with `pip install -r requirements.txt`. During development, place this repo under `ComfyUI/custom_nodes/ComfyUI_QwenVL` or symlink it there so ComfyUI autoloads the nodes. Use `python -m unittest tests.test_nodes` for the lightweight lifecycle tests, or `pytest tests -k nodes` if you prefer richer failure output. When hacking on scripts directly, export `PYTHONPATH=.` to resolve in-repo imports and set `HF_HOME` if you need a custom cache.
|
||||
|
||||
## Coding Style & Naming Conventions
|
||||
Follow PEP 8 with 4-space indentation, snake_case functions (`_clear_cuda_memory`) and PascalCase node classes (`QwenVL`, `Qwen`). Keep node attributes explicit (`RETURN_TYPES`, `CATEGORY`) so ComfyUI can surface them. Guard optional imports (e.g., `comfy.model_management`) with try/except and release GPU resources via `_maybe_move_to_cpu` before clearing CUDA. Prefer descriptive parameter names over abbreviations and document non-obvious device-handling decisions with short comments.
|
||||
|
||||
## Testing Guidelines
|
||||
`tests/test_nodes.py` relies on the standard `unittest` runner plus dummy processors/models, enabling CPU-only execution. Match that pattern when extending coverage: patch network calls, stub tensors with predictable shapes, and assert both the generated outputs and resource cleanup behaviors. Name files and methods `test_<behavior>` so discovery remains automatic, and skip GPU-specific cases when `torch.cuda.is_available()` is false. Run the suite locally before every PR and paste failures when hardware differences appear.
|
||||
|
||||
## Commit & Pull Request Guidelines
|
||||
History favors concise, imperative subjects (`Add VRAM cleanup helpers`, `Add new Qwen3 models`). Continue that style and add detail in the body only when necessary. Each PR should provide: (1) a summary of the node or workflow change, (2) test evidence (`python -m unittest` output), (3) notes on model downloads or VRAM impact, and (4) updated workflow JSON or screenshots when UI wiring changes. Cross-link GitHub issues and call out any configuration or dependency adjustments that downstream ComfyUI users must perform.
|
||||
@@ -4,14 +4,14 @@ This file provides guidance to Claude Code (claude.ai/code) when working with co
|
||||
|
||||
## Project Overview
|
||||
|
||||
ComfyUI_QwenVL is a custom node extension for ComfyUI that implements Qwen2.5-VL and Qwen2.5 language models. It provides vision-language capabilities for text-based and single-image queries, along with text-only generation. The project integrates with ComfyUI's node-based workflow system.
|
||||
ComfyUI_QwenVL is a custom node extension for ComfyUI that implements Qwen2.5-VL and Qwen3-VL vision-language models alongside Qwen2.5 text-only models. It provides multimodal AI capabilities for text-based and single-image queries, video processing, and text-only generation. The project integrates with ComfyUI's node-based workflow system through two main node classes.
|
||||
|
||||
## Architecture
|
||||
|
||||
The extension consists of two main node classes in `nodes.py`:
|
||||
|
||||
- **Qwen2VL**: Handles vision-language model inference with support for image and video inputs, now with Qwen3VL support
|
||||
- **Qwen2**: Handles text-only language model inference
|
||||
- **QwenVL**: Handles vision-language model inference with support for image and video inputs
|
||||
- **Qwen**: Handles text-only language model inference
|
||||
|
||||
Both classes follow the same architectural pattern:
|
||||
- Model loading with quantization support (none, 4bit, 8bit)
|
||||
@@ -36,6 +36,19 @@ cd ComfyUI_QwenVL
|
||||
pip install -r requirements.txt
|
||||
```
|
||||
|
||||
### Testing
|
||||
```bash
|
||||
# Run unit tests (requires PyTorch for full functionality)
|
||||
python -m unittest tests.test_nodes -v
|
||||
|
||||
# Tests use mocking to isolate functionality from ComfyUI runtime dependencies
|
||||
# Key test areas:
|
||||
# - Model loading/unloading behavior
|
||||
# - Memory management validation
|
||||
# - Quantization configuration
|
||||
# - Error handling scenarios
|
||||
```
|
||||
|
||||
### Dependencies Management
|
||||
```bash
|
||||
# Install/update dependencies
|
||||
@@ -43,22 +56,15 @@ pip install -r requirements.txt
|
||||
|
||||
# Key dependencies include:
|
||||
# - torch, torchvision, numpy, pillow
|
||||
# - huggingface_hub, transformers, bitsandbytes, accelerate
|
||||
# - huggingface_hub, transformers>=4.57.1, bitsandbytes, accelerate
|
||||
# - qwen-vl-utils, optimum
|
||||
# - transformers from git (main branch)
|
||||
```
|
||||
|
||||
### Testing
|
||||
```bash
|
||||
# Test models can be loaded (requires ComfyUI environment)
|
||||
# Load ComfyUI and verify nodes appear in "Comfyui_QwenVL" category
|
||||
```
|
||||
|
||||
## Model Configuration
|
||||
|
||||
### Supported Models
|
||||
- **Vision-Language**: Qwen2.5-VL-3B-Instruct, Qwen2.5-VL-7B-Instruct, Qwen3-VL-4B-Thinking, Qwen3-VL-8B-Thinking, SkyCaptioner-V1
|
||||
- **Text-Only**: Qwen2.5-3B/7B/14B/32B-Instruct
|
||||
- **Vision-Language**: Qwen2.5-VL-3B/7B-Instruct, Qwen3-VL-2B/4B/8B/32B-Instruct, Qwen3-VL-2B/4B/8B/32B-Thinking, SkyCaptioner-V1
|
||||
- **Text-Only**: Qwen2.5-3B/7B/14B/32B-Instruct, Qwen3-4B-Thinking-2507, Qwen3-4B-Instruct-2507
|
||||
|
||||
### Model Location
|
||||
Models are automatically downloaded to: `ComfyUI/models/LLM/`
|
||||
@@ -77,12 +83,16 @@ Models are automatically downloaded to: `ComfyUI/models/LLM/`
|
||||
|
||||
### File Organization
|
||||
```
|
||||
├── nodes.py # Main node implementations (Qwen2VL, Qwen2) with Qwen3VL support
|
||||
├── nodes.py # Main node implementations (QwenVL, Qwen) (~432 lines)
|
||||
├── __init__.py # Node class mappings for ComfyUI
|
||||
├── pyproject.toml # Project metadata and dependencies
|
||||
├── requirements.txt # Python dependencies
|
||||
├── README.md # User documentation
|
||||
└── workflow/ # Example ComfyUI workflows
|
||||
├── tests/test_nodes.py # Unit tests with mocking framework
|
||||
├── workflow/ # Example ComfyUI workflows
|
||||
│ ├── Qwen2VL.json # Multimodal workflow example
|
||||
│ ├── qwen25.json # Text generation workflow example
|
||||
│ └── *.png # Workflow screenshots
|
||||
└── README.md # User documentation
|
||||
```
|
||||
|
||||
### Node Implementation Pattern
|
||||
@@ -94,39 +104,65 @@ Both nodes follow ComfyUI's standard pattern:
|
||||
|
||||
## Important Implementation Details
|
||||
|
||||
### Qwen3VL Support
|
||||
- **Model Detection**: Uses `model.startswith("Qwen3")` to identify Qwen3 models
|
||||
- **Model Class**: Loads `Qwen3VLForConditionalGeneration` for Qwen3 models
|
||||
- **Repository Format**: Uses `Qwen/{model_name}` format for HuggingFace repository
|
||||
- **Backward Compatibility**: All existing Qwen2.5 and Skywork models continue to work unchanged
|
||||
### Model Detection and Loading Strategy
|
||||
The codebase uses intelligent model detection:
|
||||
```python
|
||||
# For vision-language models
|
||||
if model.startswith("Qwen3"):
|
||||
self.model = Qwen3VLForConditionalGeneration.from_pretrained(...)
|
||||
else:
|
||||
self.model = Qwen2_5_VLForConditionalGeneration.from_pretrained(...)
|
||||
```
|
||||
|
||||
### Video Processing
|
||||
- Uses FFmpeg for video frame extraction and resizing
|
||||
- Processes videos to 1fps with max dimension 256px
|
||||
- Creates temporary files in `/tmp/` with unique identifiers
|
||||
- Automatically cleans up temporary files after inference
|
||||
### Video Processing Pipeline
|
||||
- Uses FFmpeg subprocess calls for video frame extraction and resizing
|
||||
- Processes videos to 1fps with max dimension 256px for efficiency
|
||||
- Creates temporary files in `/tmp/` with UUID-based unique identifiers
|
||||
- Automatically cleans up temporary files after inference in finally blocks
|
||||
|
||||
### Memory Management
|
||||
- `keep_model_loaded` parameter controls model persistence
|
||||
- Automatic CUDA cache cleanup when unloading models
|
||||
### Memory Management Architecture
|
||||
- `keep_model_loaded` parameter controls model persistence between runs
|
||||
- Automatic CUDA cache cleanup when unloading models via `_clear_cuda_memory()`
|
||||
- Device detection for optimal dtype selection (bfloat16 vs float16)
|
||||
- ComfyUI integration: uses `comfy.model_management.soft_empty_cache()` when available
|
||||
|
||||
### Error Handling
|
||||
- Graceful handling of empty prompts
|
||||
- FFmpeg error handling for video processing
|
||||
- Model inference exception catching
|
||||
### Error Handling Strategy
|
||||
- Graceful handling of empty prompts with descriptive error messages
|
||||
- FFmpeg error handling for video processing failures
|
||||
- Model inference exception catching with error propagation
|
||||
- Resource cleanup in finally blocks to prevent memory leaks
|
||||
|
||||
## Development Notes
|
||||
|
||||
### Testing Framework
|
||||
The project uses a sophisticated testing approach:
|
||||
- **Mocking Strategy**: Mocks transformers, processors, and models for isolated testing
|
||||
- **ComfyUI Runtime Handling**: Creates dummy `folder_paths` module when ComfyUI unavailable
|
||||
- **Memory Testing**: Validates model loading/unloading behavior
|
||||
- **Temporary Directories**: Uses temp directories for isolated test environments
|
||||
- **Conditional Testing**: Skips tests if PyTorch not available
|
||||
|
||||
### Adding New Models
|
||||
1. Add model name to the model list in `INPUT_TYPES()`
|
||||
1. Add model name to the appropriate model list in `INPUT_TYPES()`
|
||||
2. Update model_id logic if needed (Qwen vs Skywork prefixes)
|
||||
3. Test model downloading and inference
|
||||
3. For Qwen3 models, ensure `model.startswith("Qwen3")` detection works
|
||||
4. Test model downloading and inference
|
||||
5. Update documentation if new model families are added
|
||||
|
||||
### ComfyUI Integration
|
||||
- Nodes are registered in `__init__.py` with class and display name mappings
|
||||
- Nodes are registered in `__init__.py` with multiple name mappings for compatibility
|
||||
- Category is set to "Comfyui_QwenVL"
|
||||
- Return type is always "STRING" for generated text
|
||||
- Optional integration with ComfyUI's model management system
|
||||
|
||||
### Dependencies
|
||||
The project requires specific versions of transformers and related libraries. Always use the provided requirements.txt for compatibility.
|
||||
### Dependencies and Compatibility
|
||||
- Requires `transformers>=4.57.1` for Qwen3VL support
|
||||
- Uses `qwen-vl-utils` for vision processing
|
||||
- BitsAndBytes for quantization support
|
||||
- Always use the provided requirements.txt for compatibility
|
||||
|
||||
### Performance Considerations
|
||||
- Video processing is resource-intensive; consider input size limitations
|
||||
- Model loading is expensive; use `keep_model_loaded=True` for repeated inference
|
||||
- Quantization significantly reduces memory usage but may affect output quality
|
||||
- CUDA memory management is critical for multi-model workflows
|
||||
@@ -1,4 +1,6 @@
|
||||
import os
|
||||
import gc
|
||||
import inspect
|
||||
import torch
|
||||
from transformers import (
|
||||
Qwen2_5_VLForConditionalGeneration,
|
||||
@@ -15,6 +17,42 @@ import folder_paths
|
||||
import subprocess
|
||||
import uuid
|
||||
|
||||
try:
|
||||
import comfy.model_management as comfy_mm
|
||||
except ImportError: # ComfyUI runtime not available during development/tests
|
||||
comfy_mm = None
|
||||
|
||||
|
||||
def _maybe_move_to_cpu(module):
|
||||
if module is None:
|
||||
return
|
||||
try:
|
||||
module.to("cpu")
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
def _clear_cuda_memory():
|
||||
gc.collect()
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
if comfy_mm is not None:
|
||||
try:
|
||||
soft_empty = getattr(comfy_mm, "soft_empty_cache", None)
|
||||
if callable(soft_empty):
|
||||
params = inspect.signature(soft_empty).parameters
|
||||
if "force" in params:
|
||||
soft_empty(force=True)
|
||||
else:
|
||||
soft_empty()
|
||||
return
|
||||
cleanup_models = getattr(comfy_mm, "cleanup_models", None)
|
||||
if callable(cleanup_models):
|
||||
cleanup_models()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
def tensor_to_pil(image_tensor, batch_index=0) -> Image:
|
||||
# Convert tensor of shape [batch, height, width, channels] at the batch_index to PIL Image
|
||||
@@ -37,6 +75,12 @@ class QwenVL:
|
||||
and torch.cuda.get_device_capability(self.device)[0] >= 8
|
||||
)
|
||||
|
||||
def _unload_resources(self):
|
||||
_maybe_move_to_cpu(self.model)
|
||||
self.model = None
|
||||
self.processor = None
|
||||
_clear_cuda_memory()
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
@@ -158,6 +202,9 @@ class QwenVL:
|
||||
quantization_config=quantization_config,
|
||||
)
|
||||
|
||||
processed_video_path = None
|
||||
result = None
|
||||
|
||||
with torch.no_grad():
|
||||
messages = [
|
||||
{
|
||||
@@ -168,53 +215,47 @@ class QwenVL:
|
||||
}
|
||||
]
|
||||
|
||||
if video_path:
|
||||
print("deal video_path", video_path)
|
||||
# 使用FFmpeg处理视频
|
||||
unique_id = uuid.uuid4().hex # 生成唯一标识符
|
||||
processed_video_path = f"/tmp/processed_video_{unique_id}.mp4" # 临时文件路径
|
||||
ffmpeg_command = [
|
||||
"ffmpeg",
|
||||
"-i", video_path,
|
||||
"-vf", "fps=1,scale='min(256,iw)':min'(256,ih)':force_original_aspect_ratio=decrease",
|
||||
"-c:v", "libx264",
|
||||
"-preset", "fast",
|
||||
"-crf", "18",
|
||||
processed_video_path
|
||||
]
|
||||
subprocess.run(ffmpeg_command, check=True)
|
||||
|
||||
# 添加处理后的视频信息到消息
|
||||
messages[0]["content"].insert(0, {
|
||||
"type": "video",
|
||||
"video": processed_video_path,
|
||||
})
|
||||
|
||||
# 处理图像输入
|
||||
else:
|
||||
print("deal image")
|
||||
pil_image = tensor_to_pil(image)
|
||||
messages[0]["content"].insert(0, {
|
||||
"type": "image",
|
||||
"image": pil_image,
|
||||
})
|
||||
|
||||
# 准备输入
|
||||
text = self.processor.apply_chat_template(
|
||||
messages, tokenize=False, add_generation_prompt=True
|
||||
)
|
||||
print("deal messages", messages)
|
||||
image_inputs, video_inputs = process_vision_info(messages)
|
||||
inputs = self.processor(
|
||||
text=[text],
|
||||
images=image_inputs,
|
||||
videos=video_inputs,
|
||||
padding=True,
|
||||
return_tensors="pt",
|
||||
).to("cuda")
|
||||
|
||||
# 推理
|
||||
try:
|
||||
if video_path:
|
||||
print("deal video_path", video_path)
|
||||
unique_id = uuid.uuid4().hex # 生成唯一标识符
|
||||
processed_video_path = f"/tmp/processed_video_{unique_id}.mp4" # 临时文件路径
|
||||
ffmpeg_command = [
|
||||
"ffmpeg",
|
||||
"-i", video_path,
|
||||
"-vf", "fps=1,scale='min(256,iw)':min'(256,ih)':force_original_aspect_ratio=decrease",
|
||||
"-c:v", "libx264",
|
||||
"-preset", "fast",
|
||||
"-crf", "18",
|
||||
processed_video_path
|
||||
]
|
||||
subprocess.run(ffmpeg_command, check=True)
|
||||
|
||||
messages[0]["content"].insert(0, {
|
||||
"type": "video",
|
||||
"video": processed_video_path,
|
||||
})
|
||||
else:
|
||||
print("deal image")
|
||||
pil_image = tensor_to_pil(image)
|
||||
messages[0]["content"].insert(0, {
|
||||
"type": "image",
|
||||
"image": pil_image,
|
||||
})
|
||||
|
||||
text = self.processor.apply_chat_template(
|
||||
messages, tokenize=False, add_generation_prompt=True
|
||||
)
|
||||
print("deal messages", messages)
|
||||
image_inputs, video_inputs = process_vision_info(messages)
|
||||
inputs = self.processor(
|
||||
text=[text],
|
||||
images=image_inputs,
|
||||
videos=video_inputs,
|
||||
padding=True,
|
||||
return_tensors="pt",
|
||||
).to("cuda")
|
||||
|
||||
generated_ids = self.model.generate(**inputs, max_new_tokens=max_new_tokens)
|
||||
generated_ids_trimmed = [
|
||||
out_ids[len(in_ids):] for in_ids, out_ids in zip(inputs.input_ids, generated_ids)
|
||||
@@ -227,18 +268,14 @@ class QwenVL:
|
||||
)
|
||||
except Exception as e:
|
||||
return (f"Error during model inference: {str(e)}",)
|
||||
|
||||
if not keep_model_loaded:
|
||||
del self.processor
|
||||
del self.model
|
||||
self.processor = None
|
||||
self.model = None
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
|
||||
# 删除临时视频文件
|
||||
if video_path:
|
||||
os.remove(processed_video_path)
|
||||
finally:
|
||||
if not keep_model_loaded:
|
||||
self._unload_resources()
|
||||
if processed_video_path:
|
||||
try:
|
||||
os.remove(processed_video_path)
|
||||
except FileNotFoundError:
|
||||
pass
|
||||
|
||||
return result
|
||||
|
||||
@@ -256,6 +293,12 @@ class Qwen:
|
||||
and torch.cuda.get_device_capability(self.device)[0] >= 8
|
||||
)
|
||||
|
||||
def _unload_resources(self):
|
||||
_maybe_move_to_cpu(self.model)
|
||||
self.model = None
|
||||
self.tokenizer = None
|
||||
_clear_cuda_memory()
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
@@ -354,36 +397,35 @@ class Qwen:
|
||||
quantization_config=quantization_config,
|
||||
)
|
||||
|
||||
result = None
|
||||
with torch.no_grad():
|
||||
messages = [
|
||||
{"role": "system", "content": system},
|
||||
{"role": "user", "content": prompt},
|
||||
]
|
||||
|
||||
text = self.tokenizer.apply_chat_template(
|
||||
messages, tokenize=False, add_generation_prompt=True
|
||||
)
|
||||
try:
|
||||
text = self.tokenizer.apply_chat_template(
|
||||
messages, tokenize=False, add_generation_prompt=True
|
||||
)
|
||||
|
||||
inputs = self.tokenizer([text], return_tensors="pt").to("cuda")
|
||||
inputs = self.tokenizer([text], return_tensors="pt").to("cuda")
|
||||
|
||||
generated_ids = self.model.generate(**inputs, max_new_tokens=max_new_tokens)
|
||||
generated_ids_trimmed = [
|
||||
out_ids[len(in_ids) :]
|
||||
for in_ids, out_ids in zip(inputs.input_ids, generated_ids)
|
||||
]
|
||||
result = self.tokenizer.batch_decode(
|
||||
generated_ids_trimmed,
|
||||
skip_special_tokens=True,
|
||||
clean_up_tokenization_spaces=False,
|
||||
temperature=temperature,
|
||||
)
|
||||
|
||||
if not keep_model_loaded:
|
||||
del self.tokenizer
|
||||
del self.model
|
||||
self.tokenizer = None
|
||||
self.model = None
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
generated_ids = self.model.generate(**inputs, max_new_tokens=max_new_tokens)
|
||||
generated_ids_trimmed = [
|
||||
out_ids[len(in_ids) :]
|
||||
for in_ids, out_ids in zip(inputs.input_ids, generated_ids)
|
||||
]
|
||||
result = self.tokenizer.batch_decode(
|
||||
generated_ids_trimmed,
|
||||
skip_special_tokens=True,
|
||||
clean_up_tokenization_spaces=False,
|
||||
temperature=temperature,
|
||||
)
|
||||
except Exception as e:
|
||||
return (f"Error during model inference: {str(e)}",)
|
||||
finally:
|
||||
if not keep_model_loaded:
|
||||
self._unload_resources()
|
||||
|
||||
return result
|
||||
|
||||
@@ -0,0 +1,160 @@
|
||||
import atexit
|
||||
import importlib
|
||||
import os
|
||||
import sys
|
||||
import tempfile
|
||||
import types
|
||||
import unittest
|
||||
from unittest.mock import patch
|
||||
|
||||
try:
|
||||
import torch
|
||||
except ModuleNotFoundError: # pragma: no cover - handled by skip below
|
||||
torch = None
|
||||
|
||||
|
||||
if torch is not None:
|
||||
_TEMP_DIR = tempfile.TemporaryDirectory()
|
||||
atexit.register(_TEMP_DIR.cleanup)
|
||||
folder_paths_module = types.ModuleType("folder_paths")
|
||||
folder_paths_module.models_dir = _TEMP_DIR.name
|
||||
sys.modules["folder_paths"] = folder_paths_module
|
||||
|
||||
nodes = importlib.import_module("nodes")
|
||||
nodes = importlib.reload(nodes)
|
||||
|
||||
class DummyInputs(dict):
|
||||
def __init__(self):
|
||||
tensor = torch.tensor([[0, 1]])
|
||||
super().__init__({"input_ids": tensor})
|
||||
self.input_ids = tensor
|
||||
|
||||
def to(self, device):
|
||||
self.device = device
|
||||
return self
|
||||
|
||||
class DummyProcessor:
|
||||
def apply_chat_template(self, messages, tokenize=False, add_generation_prompt=True):
|
||||
return "template"
|
||||
|
||||
def __call__(self, **kwargs):
|
||||
return DummyInputs()
|
||||
|
||||
def batch_decode(self, generated_ids_trimmed, skip_special_tokens=True, clean_up_tokenization_spaces=False, temperature=0.0):
|
||||
return ["decoded vision output"]
|
||||
|
||||
class DummyTokenizer:
|
||||
def apply_chat_template(self, messages, tokenize=False, add_generation_prompt=True):
|
||||
return "tokenized"
|
||||
|
||||
def __call__(self, texts, return_tensors="pt"):
|
||||
return DummyInputs()
|
||||
|
||||
def batch_decode(self, generated_ids_trimmed, skip_special_tokens=True, clean_up_tokenization_spaces=False, temperature=0.0):
|
||||
return ["decoded text output"]
|
||||
|
||||
class DummyVLModel:
|
||||
def __init__(self):
|
||||
self.to_calls = []
|
||||
|
||||
def to(self, device):
|
||||
self.to_calls.append(device)
|
||||
return self
|
||||
|
||||
def generate(self, **kwargs):
|
||||
return torch.tensor([[0, 1, 2, 3]])
|
||||
|
||||
class DummyTextModel(DummyVLModel):
|
||||
pass
|
||||
|
||||
class NodesTestCase(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.llm_dir = os.path.join(nodes.folder_paths.models_dir, "LLM")
|
||||
os.makedirs(self.llm_dir, exist_ok=True)
|
||||
|
||||
def _ensure_checkpoint_path(self, model_name):
|
||||
checkpoint_dir = os.path.join(self.llm_dir, model_name)
|
||||
os.makedirs(checkpoint_dir, exist_ok=True)
|
||||
return checkpoint_dir
|
||||
|
||||
@patch("nodes.process_vision_info", return_value=([None], None))
|
||||
@patch("nodes.Qwen3VLForConditionalGeneration.from_pretrained", return_value=DummyVLModel())
|
||||
@patch("nodes.AutoProcessor.from_pretrained", return_value=DummyProcessor())
|
||||
@patch("nodes._clear_cuda_memory")
|
||||
def test_qwenvl_unloads_models_after_run(self, clear_mock, _processor_mock, _model_mock, _vision_mock):
|
||||
self._ensure_checkpoint_path("Qwen3-VL-4B-Instruct")
|
||||
node = nodes.QwenVL()
|
||||
image = torch.zeros((1, 1, 1, 3))
|
||||
|
||||
result = node.inference(
|
||||
text="hi",
|
||||
model="Qwen3-VL-4B-Instruct",
|
||||
quantization="none",
|
||||
keep_model_loaded=False,
|
||||
temperature=0.7,
|
||||
max_new_tokens=10,
|
||||
seed=-1,
|
||||
image=image,
|
||||
video_path="",
|
||||
)
|
||||
|
||||
self.assertEqual(result, ["decoded vision output"])
|
||||
self.assertIsNone(node.model)
|
||||
self.assertIsNone(node.processor)
|
||||
clear_mock.assert_called_once()
|
||||
|
||||
@patch("nodes.AutoModelForCausalLM.from_pretrained", return_value=DummyTextModel())
|
||||
@patch("nodes.AutoTokenizer.from_pretrained", return_value=DummyTokenizer())
|
||||
@patch("nodes._clear_cuda_memory")
|
||||
def test_qwen_text_node_unloads_when_not_kept(self, clear_mock, _tokenizer_mock, _model_mock):
|
||||
self._ensure_checkpoint_path("Qwen3-4B-Instruct-2507")
|
||||
node = nodes.Qwen()
|
||||
|
||||
result = node.inference(
|
||||
system="sys",
|
||||
prompt="hi",
|
||||
model="Qwen3-4B-Instruct-2507",
|
||||
quantization="none",
|
||||
keep_model_loaded=False,
|
||||
temperature=0.7,
|
||||
max_new_tokens=10,
|
||||
seed=-1,
|
||||
)
|
||||
|
||||
self.assertEqual(result, ["decoded text output"])
|
||||
self.assertIsNone(node.model)
|
||||
self.assertIsNone(node.tokenizer)
|
||||
clear_mock.assert_called_once()
|
||||
|
||||
@patch("nodes.AutoModelForCausalLM.from_pretrained", return_value=DummyTextModel())
|
||||
@patch("nodes.AutoTokenizer.from_pretrained", return_value=DummyTokenizer())
|
||||
@patch("nodes._clear_cuda_memory")
|
||||
def test_qwen_text_node_keeps_model_when_requested(self, clear_mock, _tokenizer_mock, _model_mock):
|
||||
self._ensure_checkpoint_path("Qwen3-4B-Instruct-2507")
|
||||
node = nodes.Qwen()
|
||||
|
||||
node.inference(
|
||||
system="sys",
|
||||
prompt="hi",
|
||||
model="Qwen3-4B-Instruct-2507",
|
||||
quantization="none",
|
||||
keep_model_loaded=True,
|
||||
temperature=0.7,
|
||||
max_new_tokens=10,
|
||||
seed=-1,
|
||||
)
|
||||
|
||||
self.assertIsNotNone(node.model)
|
||||
self.assertIsNotNone(node.tokenizer)
|
||||
clear_mock.assert_not_called()
|
||||
|
||||
|
||||
else:
|
||||
|
||||
class NodesTestCase(unittest.TestCase):
|
||||
def test_pytorch_dependency_required(self):
|
||||
self.skipTest("PyTorch is not installed; skipping nodes tests.")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user