Add files via upload
This commit is contained in:
@@ -0,0 +1,47 @@
|
||||
# OmniGen Updates
|
||||
|
||||
## 2024-11-11
|
||||
### Added
|
||||
- Preset Prompts Support
|
||||
- Added preset prompt selection from data.json
|
||||
- Custom prompts take precedence over presets
|
||||
- Easy to extend with new presets
|
||||
- Default to empty if no preset selected
|
||||
|
||||
- Model Precision Selection
|
||||
- Added three precision options:
|
||||
- Auto: Automatically selects based on VRAM
|
||||
- FP16: Full precision (15.5GB VRAM)
|
||||
- FP8: Reduced precision (3.4GB VRAM)
|
||||
- Auto mode selects FP8 for systems with <8GB VRAM
|
||||
- Shows available VRAM in selection message
|
||||
- Smart switching between models with proper cleanup
|
||||
|
||||
- Memory Management Improvements
|
||||
- Three modes available:
|
||||
- Balanced (Default): Standard operation mode
|
||||
- Speed Priority: Keeps model in VRAM for faster consecutive generations
|
||||
- Memory Priority: Aggressive memory saving with model offloading
|
||||
- Smart model instance caching
|
||||
- Automatic VRAM cleanup when switching models
|
||||
- Recommended settings:
|
||||
- FP8 model (3.4GB VRAM): Speed Priority mode is safe
|
||||
- FP16 model (15.5GB VRAM): Use Memory Priority mode if VRAM limited
|
||||
|
||||
- Pipeline Improvements
|
||||
- Better device movement handling
|
||||
- Original pipeline backup for device movement failures
|
||||
- Improved error handling and recovery
|
||||
- Component-wise device movement for better stability
|
||||
|
||||
### Fixed
|
||||
- Pipeline device movement issues
|
||||
- Memory leaks in consecutive generations
|
||||
- Temporary file cleanup
|
||||
- Model precision switching issues
|
||||
|
||||
### Improved
|
||||
- Error handling and logging
|
||||
- VRAM usage monitoring
|
||||
- Temporary file management with UUID
|
||||
- Code organization and documentation
|
||||
+40
-67
@@ -1,76 +1,49 @@
|
||||
from typing import Dict, List, Any
|
||||
import importlib.util
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
import sys
|
||||
import os
|
||||
import importlib.util
|
||||
|
||||
# Type definitions
|
||||
NodeClassDict = Dict[str, Any]
|
||||
NodeDisplayDict = Dict[str, str]
|
||||
|
||||
# Global variables
|
||||
NODE_CLASS_MAPPINGS: NodeClassDict = {}
|
||||
NODE_DISPLAY_NAME_MAPPINGS: NodeDisplayDict = {}
|
||||
WEB_DIRECTORY = "./web"
|
||||
|
||||
def load_module(file_path: Path, module_name: str) -> None:
|
||||
"""Load a single Python module and update mappings"""
|
||||
try:
|
||||
spec = importlib.util.spec_from_file_location(module_name, file_path)
|
||||
if spec is None or spec.loader is None:
|
||||
return
|
||||
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
sys.modules[module_name] = module
|
||||
spec.loader.exec_module(module)
|
||||
|
||||
# Update mappings if they exist
|
||||
if hasattr(module, "NODE_CLASS_MAPPINGS"):
|
||||
NODE_CLASS_MAPPINGS.update(module.NODE_CLASS_MAPPINGS)
|
||||
if hasattr(module, "NODE_DISPLAY_NAME_MAPPINGS"):
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(module.NODE_DISPLAY_NAME_MAPPINGS)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
def load_modules_from_directory(directory: Path) -> None:
|
||||
"""Load all Python modules from specified directory"""
|
||||
if not directory.exists() or not directory.is_dir():
|
||||
return
|
||||
|
||||
for file_path in directory.glob("*.py"):
|
||||
if file_path.name != "__init__.py":
|
||||
load_module(file_path, file_path.stem)
|
||||
|
||||
def sort_mappings() -> None:
|
||||
"""Sort NODE_CLASS_MAPPINGS and NODE_DISPLAY_NAME_MAPPINGS"""
|
||||
global NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
|
||||
|
||||
NODE_CLASS_MAPPINGS = dict(sorted(
|
||||
NODE_CLASS_MAPPINGS.items(),
|
||||
key=lambda x: NODE_DISPLAY_NAME_MAPPINGS.get(x[0], x[0])
|
||||
))
|
||||
NODE_DISPLAY_NAME_MAPPINGS = dict(sorted(
|
||||
NODE_DISPLAY_NAME_MAPPINGS.items(),
|
||||
key=lambda x: x[1]
|
||||
))
|
||||
|
||||
def load_javascript(web_directory: str) -> List[Dict[str, str]]:
|
||||
"""Return JavaScript file configurations"""
|
||||
return [{"path": "refreshNode.js"}]
|
||||
|
||||
# Initialize module loading
|
||||
# Add module directory to Python path
|
||||
current_dir = Path(__file__).parent
|
||||
sys.path.insert(0, str(current_dir))
|
||||
if str(current_dir) not in sys.path:
|
||||
sys.modules[__name__] = sys.modules.get(__name__, type(__name__, (), {}))
|
||||
sys.path.insert(0, str(current_dir))
|
||||
|
||||
# Load modules from current directory and py subdirectory
|
||||
load_modules_from_directory(current_dir)
|
||||
load_modules_from_directory(current_dir / "py")
|
||||
sort_mappings()
|
||||
# Initialize mappings
|
||||
NODE_CLASS_MAPPINGS = {}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {}
|
||||
|
||||
def load_nodes():
|
||||
"""Automatically discover and load node definitions"""
|
||||
for file in current_dir.glob("*.py"):
|
||||
if file.stem == "__init__":
|
||||
continue
|
||||
|
||||
try:
|
||||
# Import module
|
||||
spec = importlib.util.spec_from_file_location(file.stem, file)
|
||||
if spec and spec.loader:
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
spec.loader.exec_module(module)
|
||||
|
||||
# Update mappings
|
||||
if hasattr(module, "NODE_CLASS_MAPPINGS"):
|
||||
NODE_CLASS_MAPPINGS.update(module.NODE_CLASS_MAPPINGS)
|
||||
if hasattr(module, "NODE_DISPLAY_NAME_MAPPINGS"):
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(module.NODE_DISPLAY_NAME_MAPPINGS)
|
||||
|
||||
# Initialize paths if available
|
||||
if hasattr(module, "Paths") and hasattr(module.Paths, "LLM_DIR"):
|
||||
os.makedirs(module.Paths.LLM_DIR, exist_ok=True)
|
||||
|
||||
except Exception as e:
|
||||
print(f"Error loading {file.name}: {str(e)}")
|
||||
|
||||
# Load all nodes
|
||||
load_nodes()
|
||||
|
||||
__all__ = [
|
||||
"NODE_CLASS_MAPPINGS",
|
||||
"NODE_DISPLAY_NAME_MAPPINGS",
|
||||
"WEB_DIRECTORY",
|
||||
"load_javascript"
|
||||
"NODE_DISPLAY_NAME_MAPPINGS"
|
||||
]
|
||||
|
||||
|
||||
+401
-240
@@ -1,240 +1,401 @@
|
||||
import sys
|
||||
import os.path as osp
|
||||
import os
|
||||
import torch
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
from huggingface_hub import snapshot_download
|
||||
import requests
|
||||
import folder_paths
|
||||
import tempfile
|
||||
import shutil
|
||||
|
||||
# Define all path constants
|
||||
class Paths:
|
||||
ROOT_DIR = osp.dirname(__file__)
|
||||
MODELS_DIR = folder_paths.models_dir
|
||||
LLM_DIR = osp.join(MODELS_DIR, "LLM")
|
||||
OMNIGEN_DIR = osp.join(LLM_DIR, "OmniGen-v1")
|
||||
OMNIGEN_CODE_DIR = osp.join(ROOT_DIR, "OmniGen")
|
||||
TMP_DIR = osp.join(ROOT_DIR, "tmp")
|
||||
MODEL_FILE = osp.join(OMNIGEN_DIR, "model.safetensors")
|
||||
|
||||
# Ensure necessary directories exist
|
||||
os.makedirs(Paths.LLM_DIR, exist_ok=True)
|
||||
sys.path.append(Paths.ROOT_DIR)
|
||||
|
||||
class ailab_OmniGen:
|
||||
def __init__(self):
|
||||
self._ensure_code_exists()
|
||||
self._ensure_model_exists()
|
||||
try:
|
||||
from OmniGen import OmniGenPipeline
|
||||
self.OmniGenPipeline = OmniGenPipeline
|
||||
except ImportError as e:
|
||||
print(f"Error importing OmniGen: {e}")
|
||||
raise RuntimeError("Failed to import OmniGen. Please check if the code was downloaded correctly.")
|
||||
|
||||
def _ensure_code_exists(self):
|
||||
"""Ensure OmniGen code exists, download from Hugging Face if not"""
|
||||
try:
|
||||
if not osp.exists(Paths.OMNIGEN_CODE_DIR):
|
||||
print("Downloading OmniGen code from Hugging Face...")
|
||||
|
||||
# Files to download from Hugging Face
|
||||
files = [
|
||||
"model.py",
|
||||
"pipeline.py",
|
||||
"processor.py",
|
||||
"scheduler.py",
|
||||
"transformer.py",
|
||||
"utils.py",
|
||||
"__init__.py"
|
||||
]
|
||||
|
||||
os.makedirs(Paths.OMNIGEN_CODE_DIR, exist_ok=True)
|
||||
base_url = "https://huggingface.co/spaces/Shitao/OmniGen/raw/main/OmniGen/"
|
||||
|
||||
for file in files:
|
||||
url = base_url + file
|
||||
response = requests.get(url)
|
||||
if response.status_code == 200:
|
||||
with open(osp.join(Paths.OMNIGEN_CODE_DIR, file), 'wb') as f:
|
||||
f.write(response.content)
|
||||
print(f"Downloaded {file}")
|
||||
else:
|
||||
raise RuntimeError(f"Failed to download {file}: {response.status_code}")
|
||||
|
||||
print("OmniGen code setup completed")
|
||||
|
||||
if Paths.OMNIGEN_CODE_DIR not in sys.path:
|
||||
sys.path.append(Paths.OMNIGEN_CODE_DIR)
|
||||
|
||||
else:
|
||||
print("OmniGen code already exists")
|
||||
|
||||
except Exception as e:
|
||||
print(f"Error downloading OmniGen code: {e}")
|
||||
raise RuntimeError(f"Failed to download OmniGen code: {str(e)}")
|
||||
|
||||
def _ensure_model_exists(self):
|
||||
"""Ensure model file exists, download if not"""
|
||||
try:
|
||||
if not osp.exists(Paths.MODEL_FILE):
|
||||
print("OmniGen model not found, starting download from Hugging Face...")
|
||||
os.makedirs(Paths.OMNIGEN_DIR, exist_ok=True)
|
||||
snapshot_download(
|
||||
repo_id="Shitao/OmniGen-v1",
|
||||
local_dir=Paths.OMNIGEN_DIR,
|
||||
local_dir_use_symlinks=False,
|
||||
resume_download=True,
|
||||
token=None, # Add your token if needed
|
||||
tqdm_class=None, # This will use default progress bar
|
||||
)
|
||||
print("OmniGen model downloaded successfully")
|
||||
else:
|
||||
print("OmniGen model found locally")
|
||||
except Exception as e:
|
||||
print(f"Error during model initialization: {e}")
|
||||
raise RuntimeError(f"Failed to initialize OmniGen model: {str(e)}")
|
||||
|
||||
def _setup_temp_dir(self):
|
||||
"""Set up temporary directory"""
|
||||
if osp.exists(Paths.TMP_DIR):
|
||||
shutil.rmtree(Paths.TMP_DIR)
|
||||
os.makedirs(Paths.TMP_DIR, exist_ok=True)
|
||||
|
||||
def _cleanup_temp_dir(self):
|
||||
"""Clean up temporary directory"""
|
||||
if osp.exists(Paths.TMP_DIR):
|
||||
shutil.rmtree(Paths.TMP_DIR)
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"prompt": ("STRING", {"multiline": True, "forceInput": False, "default": ""}),
|
||||
"guidance_scale": ("FLOAT", {"default": 3.5, "min": 1.0, "max": 5.0, "step": 0.1, "round": 0.01}),
|
||||
"img_guidance_scale": ("FLOAT", {"default": 1.8, "min": 1.0, "max": 2.0, "step": 0.1, "round": 0.01}),
|
||||
"num_inference_steps": ("INT", {"default": 50, "min": 1, "max": 100, "step": 1}),
|
||||
"separate_cfg_infer": ("BOOLEAN", {"default": True}),
|
||||
"offload_model": ("BOOLEAN", {"default": False}),
|
||||
"use_input_image_size_as_output": ("BOOLEAN", {"default": False}),
|
||||
"width": ("INT", {"default": 1024, "min": 128, "max": 2048, "step": 16}),
|
||||
"height": ("INT", {"default": 1024, "min": 128, "max": 2048, "step": 16}),
|
||||
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
|
||||
"max_input_image_size": ("INT", {"default": 1024, "min": 128, "max": 2048, "step": 16}),
|
||||
},
|
||||
"optional": {
|
||||
"image_1": ("IMAGE",),
|
||||
"image_2": ("IMAGE",),
|
||||
"image_3": ("IMAGE",),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
FUNCTION = "generation"
|
||||
CATEGORY = "🧪AILab/OmniGen"
|
||||
|
||||
def save_input_img(self, image):
|
||||
try:
|
||||
with tempfile.NamedTemporaryFile(suffix=".png", delete=False, dir=Paths.TMP_DIR) as f:
|
||||
img_np = image.numpy()[0] * 255
|
||||
img_pil = Image.fromarray(img_np.astype(np.uint8))
|
||||
img_pil.save(f.name)
|
||||
return f.name
|
||||
except Exception as e:
|
||||
print(f"Error saving input image: {e}")
|
||||
return None
|
||||
|
||||
def _process_prompt_and_images(self, prompt, images):
|
||||
"""Process prompt and images, return updated prompt and image paths"""
|
||||
input_images = []
|
||||
|
||||
# Auto-generate prompt if empty but images provided
|
||||
if not prompt and any(images):
|
||||
prompt = " ".join(f"<img><|image_{i+1}|></img>" for i, img in enumerate(images) if img is not None)
|
||||
|
||||
# Process each image
|
||||
for i, img in enumerate(images, 1):
|
||||
if img is not None:
|
||||
img_path = self.save_input_img(img)
|
||||
if img_path:
|
||||
input_images.append(img_path)
|
||||
img_tag = f"<img><|image_{i}|></img>"
|
||||
# Handle both image_1 and image1 formats
|
||||
if f"image_{i}" in prompt:
|
||||
prompt = prompt.replace(f"image_{i}", img_tag)
|
||||
elif f"image{i}" in prompt: # Added support for image1 format
|
||||
prompt = prompt.replace(f"image{i}", img_tag)
|
||||
elif img_tag not in prompt:
|
||||
prompt += f" {img_tag}"
|
||||
|
||||
return prompt, input_images
|
||||
|
||||
def generation(self, prompt, num_inference_steps, guidance_scale,
|
||||
img_guidance_scale, max_input_image_size, separate_cfg_infer, offload_model,
|
||||
use_input_image_size_as_output, width, height, seed, image_1=None, image_2=None, image_3=None):
|
||||
try:
|
||||
self._setup_temp_dir()
|
||||
pipe = self.OmniGenPipeline.from_pretrained(Paths.OMNIGEN_DIR)
|
||||
|
||||
# Switch to eager mode only if SDPA is not supported
|
||||
if not self._check_sdpa_support():
|
||||
if hasattr(pipe, 'text_encoder'):
|
||||
pipe.text_encoder.config.attn_implementation = "eager"
|
||||
if hasattr(pipe, 'unet'):
|
||||
pipe.unet.config.attn_implementation = "eager"
|
||||
|
||||
# Process prompt and images
|
||||
prompt, input_images = self._process_prompt_and_images(prompt, [image_1, image_2, image_3])
|
||||
input_images = input_images if input_images else None
|
||||
|
||||
print(f"Processing with prompt: {prompt}")
|
||||
output = pipe(
|
||||
prompt=prompt,
|
||||
input_images=input_images,
|
||||
guidance_scale=guidance_scale,
|
||||
img_guidance_scale=img_guidance_scale,
|
||||
num_inference_steps=num_inference_steps,
|
||||
separate_cfg_infer=separate_cfg_infer,
|
||||
use_kv_cache=True,
|
||||
offload_kv_cache=True,
|
||||
offload_model=offload_model,
|
||||
use_input_image_size_as_output=use_input_image_size_as_output,
|
||||
width=width,
|
||||
height=height,
|
||||
seed=seed,
|
||||
max_input_image_size=max_input_image_size,
|
||||
)
|
||||
|
||||
img = np.array(output[0]) / 255.0
|
||||
img = torch.from_numpy(img).unsqueeze(0)
|
||||
return (img,)
|
||||
|
||||
except Exception as e:
|
||||
print(f"Error during generation: {e}")
|
||||
raise e
|
||||
finally:
|
||||
self._cleanup_temp_dir()
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
def _check_sdpa_support(self):
|
||||
"""Check if system supports SDPA"""
|
||||
try:
|
||||
import torch
|
||||
if hasattr(torch.nn.functional, 'scaled_dot_product_attention'):
|
||||
return True
|
||||
except Exception as e:
|
||||
print(f"Warning: SDPA not supported, falling back to eager attention implementation")
|
||||
|
||||
return False
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"ailab_OmniGen": ailab_OmniGen
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"ailab_OmniGen": "OmniGen 🖼️"
|
||||
}
|
||||
import sys
|
||||
import os.path as osp
|
||||
import os
|
||||
import torch
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
from huggingface_hub import snapshot_download
|
||||
import requests
|
||||
import folder_paths
|
||||
import tempfile
|
||||
import shutil
|
||||
import json
|
||||
import uuid
|
||||
|
||||
# Define all path constants
|
||||
class Paths:
|
||||
ROOT_DIR = osp.dirname(__file__)
|
||||
MODELS_DIR = folder_paths.models_dir
|
||||
LLM_DIR = osp.join(MODELS_DIR, "LLM")
|
||||
OMNIGEN_DIR = osp.join(LLM_DIR, "OmniGen-v1")
|
||||
OMNIGEN_CODE_DIR = osp.join(ROOT_DIR, "OmniGen")
|
||||
TMP_DIR = osp.join(ROOT_DIR, "tmp")
|
||||
MODEL_FILE_FP16 = osp.join(OMNIGEN_DIR, "model.safetensors")
|
||||
MODEL_FILE_FP8 = osp.join(OMNIGEN_DIR, "model-fp8_e4m3fn.safetensors")
|
||||
|
||||
# Ensure necessary directories exist
|
||||
os.makedirs(Paths.LLM_DIR, exist_ok=True)
|
||||
sys.path.append(Paths.ROOT_DIR)
|
||||
|
||||
class ailab_OmniGen:
|
||||
VERSION = "1.2.0"
|
||||
_model_instance = None
|
||||
_current_precision = None
|
||||
|
||||
# Load preset prompts
|
||||
try:
|
||||
json_path = osp.join(osp.dirname(__file__), "data.json")
|
||||
if osp.exists(json_path):
|
||||
with open(json_path, 'r', encoding='utf-8') as f:
|
||||
data = json.load(f)
|
||||
PRESET_PROMPTS = data.get("PRESET_PROMPTS", {"None": ""})
|
||||
else:
|
||||
PRESET_PROMPTS = {"None": ""}
|
||||
except Exception as e:
|
||||
print(f"Error loading preset prompts: {e}")
|
||||
PRESET_PROMPTS = {"None": ""}
|
||||
|
||||
def __init__(self):
|
||||
self._ensure_code_exists()
|
||||
self._ensure_model_exists()
|
||||
|
||||
try:
|
||||
from OmniGen import OmniGenPipeline
|
||||
self.OmniGenPipeline = OmniGenPipeline
|
||||
except ImportError as e:
|
||||
print(f"Error importing OmniGen: {e}")
|
||||
raise RuntimeError("Failed to import OmniGen. Please check if the code was downloaded correctly.")
|
||||
|
||||
def _ensure_code_exists(self):
|
||||
"""Ensure OmniGen code exists, download from Hugging Face if not"""
|
||||
try:
|
||||
if not osp.exists(Paths.OMNIGEN_CODE_DIR):
|
||||
print("Downloading OmniGen code from Hugging Face...")
|
||||
|
||||
# Files to download from Hugging Face
|
||||
files = [
|
||||
"model.py",
|
||||
"pipeline.py",
|
||||
"processor.py",
|
||||
"scheduler.py",
|
||||
"transformer.py",
|
||||
"utils.py",
|
||||
"__init__.py"
|
||||
]
|
||||
|
||||
os.makedirs(Paths.OMNIGEN_CODE_DIR, exist_ok=True)
|
||||
base_url = "https://huggingface.co/spaces/Shitao/OmniGen/raw/main/OmniGen/"
|
||||
|
||||
for file in files:
|
||||
url = base_url + file
|
||||
response = requests.get(url)
|
||||
if response.status_code == 200:
|
||||
with open(osp.join(Paths.OMNIGEN_CODE_DIR, file), 'wb') as f:
|
||||
f.write(response.content)
|
||||
print(f"Downloaded {file}")
|
||||
else:
|
||||
raise RuntimeError(f"Failed to download {file}: {response.status_code}")
|
||||
|
||||
print("OmniGen code setup completed")
|
||||
|
||||
if Paths.OMNIGEN_CODE_DIR not in sys.path:
|
||||
sys.path.append(Paths.OMNIGEN_CODE_DIR)
|
||||
|
||||
else:
|
||||
print("OmniGen code already exists")
|
||||
|
||||
except Exception as e:
|
||||
print(f"Error downloading OmniGen code: {e}")
|
||||
raise RuntimeError(f"Failed to download OmniGen code: {str(e)}")
|
||||
|
||||
def _ensure_model_exists(self, model_precision=None):
|
||||
"""Ensure model file exists, download if not"""
|
||||
try:
|
||||
os.makedirs(Paths.OMNIGEN_DIR, exist_ok=True)
|
||||
|
||||
# Download FP8 model if specified and not exists
|
||||
if model_precision == "FP8" and not osp.exists(Paths.MODEL_FILE_FP8):
|
||||
print("FP8 model not found, downloading from Hugging Face...")
|
||||
url = "https://huggingface.co/silveroxides/OmniGen-V1/resolve/main/model-fp8_e4m3fn.safetensors"
|
||||
response = requests.get(url, stream=True)
|
||||
if response.status_code == 200:
|
||||
with open(Paths.MODEL_FILE_FP8, 'wb') as f:
|
||||
for chunk in response.iter_content(chunk_size=8192):
|
||||
if chunk:
|
||||
f.write(chunk)
|
||||
print("FP8 model downloaded successfully")
|
||||
else:
|
||||
raise RuntimeError(f"Failed to download FP8 model: {response.status_code}")
|
||||
|
||||
# Check if FP16 model exists
|
||||
if not osp.exists(Paths.MODEL_FILE_FP16):
|
||||
print("FP16 model not found, starting download from Hugging Face...")
|
||||
snapshot_download(
|
||||
repo_id="silveroxides/OmniGen-V1",
|
||||
local_dir=Paths.OMNIGEN_DIR,
|
||||
local_dir_use_symlinks=False,
|
||||
resume_download=True,
|
||||
token=None,
|
||||
tqdm_class=None,
|
||||
)
|
||||
print("FP16 model downloaded successfully")
|
||||
|
||||
# Verify model files exist after download
|
||||
if model_precision == "FP8" and not osp.exists(Paths.MODEL_FILE_FP8):
|
||||
raise RuntimeError("FP8 model download failed")
|
||||
if not osp.exists(Paths.MODEL_FILE_FP16):
|
||||
raise RuntimeError("FP16 model download failed")
|
||||
|
||||
print("OmniGen models verified successfully")
|
||||
|
||||
except Exception as e:
|
||||
print(f"Error during model initialization: {e}")
|
||||
raise RuntimeError(f"Failed to initialize OmniGen model: {str(e)}")
|
||||
|
||||
def _setup_temp_dir(self):
|
||||
"""Set up temporary directory with unique name"""
|
||||
self._temp_dir = osp.join(Paths.TMP_DIR, str(uuid.uuid4()))
|
||||
os.makedirs(self._temp_dir, exist_ok=True)
|
||||
|
||||
def _cleanup_temp_dir(self):
|
||||
"""Clean up temporary directory"""
|
||||
if hasattr(self, '_temp_dir') and osp.exists(self._temp_dir):
|
||||
shutil.rmtree(self._temp_dir)
|
||||
|
||||
def _auto_select_precision(self):
|
||||
"""Automatically select precision based on available VRAM"""
|
||||
if torch.cuda.is_available():
|
||||
vram_size = torch.cuda.get_device_properties(0).total_memory / 1024**3 # GB
|
||||
if vram_size < 8: # 如果VRAM小于8GB
|
||||
print(f"Auto selecting FP8 (Available VRAM: {vram_size:.1f}GB)")
|
||||
return "FP8"
|
||||
print(f"Auto selecting FP16 (Available VRAM: {vram_size:.1f}GB)")
|
||||
return "FP16"
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"preset_prompt": (list(s.PRESET_PROMPTS.keys()), {"default": "None"}),
|
||||
"prompt": ("STRING", {"multiline": True, "forceInput": False, "default": ""}),
|
||||
"model_precision": (["Auto", "FP16", "FP8"], {"default": "Auto"}),
|
||||
"memory_management": (["Balanced", "Speed Priority", "Memory Priority"], {"default": "Balanced"}),
|
||||
"guidance_scale": ("FLOAT", {"default": 3.5, "min": 1.0, "max": 5.0, "step": 0.1, "round": 0.01}),
|
||||
"img_guidance_scale": ("FLOAT", {"default": 1.8, "min": 1.0, "max": 2.0, "step": 0.1, "round": 0.01}),
|
||||
"num_inference_steps": ("INT", {"default": 50, "min": 1, "max": 100, "step": 1}),
|
||||
"separate_cfg_infer": ("BOOLEAN", {"default": True}),
|
||||
"use_input_image_size_as_output": ("BOOLEAN", {"default": False}),
|
||||
"width": ("INT", {"default": 512, "min": 128, "max": 2048, "step": 8}),
|
||||
"height": ("INT", {"default": 512, "min": 128, "max": 2048, "step": 8}),
|
||||
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
|
||||
"max_input_image_size": ("INT", {"default": 1024, "min": 128, "max": 2048, "step": 16}),
|
||||
},
|
||||
"optional": {
|
||||
"image_1": ("IMAGE",),
|
||||
"image_2": ("IMAGE",),
|
||||
"image_3": ("IMAGE",),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
FUNCTION = "generation"
|
||||
CATEGORY = "🧪AILab/OmniGen"
|
||||
|
||||
def save_input_img(self, image):
|
||||
try:
|
||||
with tempfile.NamedTemporaryFile(suffix=".png", delete=False, dir=Paths.TMP_DIR) as f:
|
||||
img_np = image.numpy()[0] * 255
|
||||
img_pil = Image.fromarray(img_np.astype(np.uint8))
|
||||
img_pil.save(f.name)
|
||||
return f.name
|
||||
except Exception as e:
|
||||
print(f"Error saving input image: {e}")
|
||||
return None
|
||||
|
||||
def _process_prompt_and_images(self, prompt, images):
|
||||
"""Process prompt and images, return updated prompt and image paths"""
|
||||
input_images = []
|
||||
|
||||
# Auto-generate prompt if empty but images provided
|
||||
if not prompt and any(images):
|
||||
prompt = " ".join(f"<img><|image_{i+1}|></img>" for i, img in enumerate(images) if img is not None)
|
||||
|
||||
# Process each image
|
||||
for i, img in enumerate(images, 1):
|
||||
if img is not None:
|
||||
img_path = self.save_input_img(img)
|
||||
if img_path:
|
||||
input_images.append(img_path)
|
||||
img_tag = f"<img><|image_{i}|></img>"
|
||||
if f"image_{i}" in prompt:
|
||||
prompt = prompt.replace(f"image_{i}", img_tag)
|
||||
elif f"image{i}" in prompt:
|
||||
prompt = prompt.replace(f"image{i}", img_tag)
|
||||
elif img_tag not in prompt:
|
||||
prompt += f" {img_tag}"
|
||||
|
||||
return prompt.strip(), input_images
|
||||
|
||||
def _check_sdpa_support(self):
|
||||
"""Check if system supports Scaled Dot Product Attention"""
|
||||
try:
|
||||
import torch
|
||||
if hasattr(torch.nn.functional, 'scaled_dot_product_attention'):
|
||||
return True
|
||||
return False
|
||||
except Exception as e:
|
||||
print(f"Error checking SDPA support: {e}")
|
||||
return False
|
||||
|
||||
def _get_pipeline(self, model_precision, keep_in_vram):
|
||||
try:
|
||||
# Reuse existing instance if available
|
||||
if keep_in_vram and self._model_instance and self._current_precision == model_precision:
|
||||
print("Reusing existing pipeline instance")
|
||||
return self._model_instance
|
||||
|
||||
# Check model file
|
||||
model_file = Paths.MODEL_FILE_FP8 if model_precision == "FP8" else Paths.MODEL_FILE_FP16
|
||||
if not os.path.exists(model_file):
|
||||
raise RuntimeError(f"Model file not found: {model_file}")
|
||||
|
||||
# Create pipeline
|
||||
try:
|
||||
# Initialize pipeline
|
||||
pipe = self.OmniGenPipeline.from_pretrained(Paths.OMNIGEN_DIR)
|
||||
|
||||
if pipe is None:
|
||||
raise RuntimeError("Initial pipeline creation failed")
|
||||
|
||||
# Save original pipeline reference before moving to device
|
||||
original_pipe = pipe
|
||||
|
||||
# Move to device
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
try:
|
||||
# Move model components first
|
||||
if hasattr(pipe, 'text_encoder'):
|
||||
pipe.text_encoder = pipe.text_encoder.to(device)
|
||||
if hasattr(pipe, 'unet'):
|
||||
pipe.unet = pipe.unet.to(device)
|
||||
if hasattr(pipe, 'vae'):
|
||||
pipe.vae = pipe.vae.to(device)
|
||||
|
||||
# Then move entire pipeline
|
||||
pipe = pipe.to(device)
|
||||
|
||||
# Use original pipeline if None after moving
|
||||
if pipe is None:
|
||||
print("Warning: Pipeline.to(device) returned None, using original pipeline")
|
||||
pipe = original_pipe
|
||||
|
||||
except Exception as e:
|
||||
print(f"Warning: Error moving pipeline to device: {e}, using original pipeline")
|
||||
pipe = original_pipe
|
||||
|
||||
# Validate pipeline
|
||||
if not callable(pipe):
|
||||
raise RuntimeError("Pipeline is not callable after initialization")
|
||||
|
||||
# Save instance if needed
|
||||
if keep_in_vram:
|
||||
self._model_instance = pipe
|
||||
self._current_precision = model_precision
|
||||
|
||||
return pipe
|
||||
|
||||
except Exception as pipe_error:
|
||||
print(f"Pipeline creation error: {pipe_error}")
|
||||
raise
|
||||
|
||||
except Exception as e:
|
||||
print(f"Fatal error in pipeline creation: {str(e)}")
|
||||
raise RuntimeError(f"Failed to create pipeline: {str(e)}")
|
||||
|
||||
def generation(self, preset_prompt, model_precision, prompt, memory_management, num_inference_steps, guidance_scale,
|
||||
img_guidance_scale, max_input_image_size, separate_cfg_infer,
|
||||
use_input_image_size_as_output, width, height, seed,
|
||||
image_1=None, image_2=None, image_3=None):
|
||||
try:
|
||||
# 如果选择Auto,自动选择精度
|
||||
if model_precision == "Auto":
|
||||
model_precision = self._auto_select_precision()
|
||||
|
||||
self._setup_temp_dir()
|
||||
|
||||
# 清理现有实例如果精度不匹配
|
||||
if self._model_instance and self._current_precision != model_precision:
|
||||
print(f"Precision changed from {self._current_precision} to {model_precision}, clearing instance")
|
||||
self._model_instance = None
|
||||
self._current_precision = None
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
print("VRAM cleared")
|
||||
|
||||
# 内存管理策略
|
||||
if memory_management == "Memory Priority":
|
||||
print("Memory Priority mode: Forcing pipeline recreation")
|
||||
self._model_instance = None
|
||||
self._current_precision = None
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
keep_in_vram = (memory_management == "Speed Priority")
|
||||
offload_model = (memory_management == "Memory Priority")
|
||||
|
||||
# Check model instance status
|
||||
print(f"Current model instance: {'Present' if self._model_instance else 'None'}")
|
||||
print(f"Current model precision: {self._current_precision}")
|
||||
|
||||
final_prompt = prompt.strip() if prompt.strip() else self.PRESET_PROMPTS[preset_prompt]
|
||||
pipe = self._get_pipeline(model_precision, keep_in_vram)
|
||||
|
||||
# Monitor VRAM usage
|
||||
if torch.cuda.is_available():
|
||||
print(f"VRAM usage after pipeline creation: {torch.cuda.memory_allocated()/1024**2:.2f}MB")
|
||||
|
||||
# Process prompt and images
|
||||
final_prompt, input_images = self._process_prompt_and_images(final_prompt, [image_1, image_2, image_3])
|
||||
input_images = input_images if input_images else None
|
||||
|
||||
print(f"Processing with prompt: {final_prompt}")
|
||||
print(f"Model will be {'offloaded' if offload_model else 'kept'} during generation")
|
||||
|
||||
output = pipe(
|
||||
prompt=final_prompt,
|
||||
input_images=input_images,
|
||||
guidance_scale=guidance_scale,
|
||||
img_guidance_scale=img_guidance_scale,
|
||||
num_inference_steps=num_inference_steps,
|
||||
separate_cfg_infer=separate_cfg_infer,
|
||||
use_kv_cache=True,
|
||||
offload_kv_cache=True,
|
||||
offload_model=offload_model,
|
||||
use_input_image_size_as_output=use_input_image_size_as_output,
|
||||
width=width,
|
||||
height=height,
|
||||
seed=seed,
|
||||
max_input_image_size=max_input_image_size,
|
||||
)
|
||||
|
||||
# Print VRAM usage after generation
|
||||
if torch.cuda.is_available():
|
||||
print(f"VRAM usage after generation: {torch.cuda.memory_allocated()/1024**2:.2f}MB")
|
||||
|
||||
img = np.array(output[0]) / 255.0
|
||||
img = torch.from_numpy(img).unsqueeze(0)
|
||||
|
||||
# Clean up if not keeping in VRAM
|
||||
if not keep_in_vram:
|
||||
del pipe
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
return (img,)
|
||||
|
||||
except Exception as e:
|
||||
print(f"Error during generation: {e}")
|
||||
raise e
|
||||
finally:
|
||||
self._cleanup_temp_dir()
|
||||
if not keep_in_vram and torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"ailab_OmniGen": ailab_OmniGen
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"ailab_OmniGen": "OmniGen 🖼️"
|
||||
}
|
||||
@@ -0,0 +1,35 @@
|
||||
{
|
||||
"PRESET_PROMPTS": {
|
||||
"None": "",
|
||||
"20yo woman looking at viewer": "Create an image of a 20-year-old woman looking directly at the viewer, with a neutral or friendly expression.",
|
||||
"Transform image_1 into an oil painting (image_1)": "Transform image_1 into an oil painting, giving it a textured, classic style with visible brushstrokes and rich color.",
|
||||
"Transform image_1 into an Anime (image_1)": "Transform image_1 into an anime-style illustration, with large expressive eyes, vibrant colors, and exaggerated features.",
|
||||
"The girl in image_1 sitting on rock on top of the mountain (image_1)": "Depict the girl from image_1 sitting on a rock at the top of a mountain, gazing out over a breathtaking landscape.",
|
||||
"Combine 2 People in anime style (image_1, image_2)": "Combine the characters from image_1 and image_2 in anime style, blending their features and surroundings into one cohesive scene.",
|
||||
"2 people at the coffee shop (image_1, image_2)": "A woman from image_1 and a man from image_2 are sitting across from each other at a cozy coffee shop, each holding a cup of coffee and engaging in conversation.",
|
||||
"Depth map to image (image_1)": "Following the depth mapping of image_1, generate a new photo: an elderly couple sitting at a cozy coffee shop, with layers of depth and focus.",
|
||||
"Image to pose skeleton (image_1)": "Detect the skeleton of a human in image_1, creating a skeletal overlay for analysis or artistic interpretation.",
|
||||
"Pose skeleton to image (image_1)": "Following the pose of the human skeleton detected in image_1, generate a new photo of the subject in the same pose with realistic anatomy.",
|
||||
"Deblur image (image_1)": "Deblur this image: image_1, removing any motion blur or focus issues to create a clearer, more defined image.",
|
||||
"Make an object come to life (image_1)": "Turn, an inanimate object in image_1 like a teapot, into a lively character with eyes, a smile, and moving limbs.",
|
||||
"Transform a landscape (image_1)": "Transform the serene mountain landscape in image_1 into a glowing, magical world with floating islands and sparkling rivers.",
|
||||
"Mix people and background (image_1, image_2)": "Use image_1 as the background and place the girl wearing a red dress from image_2 in the foreground, making her appear as if she’s walking through a foggy forest.",
|
||||
"Combine creatures (image_1, image_2)": "Combine the fierce lion from image_1 with the majestic eagle in image_2 to create a mythical creature with the body of a lion and wings of an eagle.",
|
||||
"Create a futuristic city (image_1)": "Use image_1, a city skyline, and transform it into a futuristic metropolis with flying cars, neon lights, and holograms.",
|
||||
"Fantasy world building (image_1, image_2, image_3)": "Mix the icy landscape from image_1, the mystical castle from image_2, and the dark forest from image_3 to build a fantasy world full of adventure.",
|
||||
"Create a weather transformation (image_1)": "Take the sunny day in image_1 and transform it into a dramatic thunderstorm with dark clouds, lightning strikes, and strong winds.",
|
||||
"Surreal composition (image_1)": "Create a surreal scene by blending image_1, where the sky is full of floating clocks, melting trees, and a river made of clouds.",
|
||||
"Turn a person into a mythical creature (image_1)": "Take the portrait of the person in image_1 and transform them into a beautiful, ethereal creature with wings, glowing eyes, and a radiant aura.",
|
||||
"Underwater scene (image_1)": "Place the subject of image_1 underwater, surrounded by schools of fish, vibrant coral reefs, and beams of sunlight filtering through the water.",
|
||||
"Create a time-lapse effect (image_1)": "Turn image_1 into a time-lapse scene where flowers bloom, the sun rises and sets, and a bustling city grows in the background.",
|
||||
"Epic battle scene (image_1, image_2)": "Combine the knight in image_1 with the dragon in image_2 to create an epic battle scene in a fiery wasteland.",
|
||||
"Create a dreamlike atmosphere (image_1)": "Use image_1, a forest, and turn it into a dreamlike scene where the trees glow, the ground sparkles, and the sky is filled with colorful swirls of light.",
|
||||
"Mix old and new (image_1, image_2)": "Combine the vintage car from image_1 with the futuristic cityscape from image_2, making the car drive through a neon-lit street full of towering skyscrapers.",
|
||||
"Turn a mundane object into art (image_1)": "Take the everyday coffee mug from image_1 and turn it into a masterpiece, where the mug is surrounded by swirling colors, abstract shapes, and vibrant patterns.",
|
||||
"Create a superhero scene (image_1, image_2)": "Use image_1 as the setting, and place the superhero from image_2 in the center of the city, with lightning and energy blasts emanating from their hands.",
|
||||
"Make a monster (image_1)": "Transform the image of a simple animal in image_1 into a terrifying, mythical monster with glowing eyes, sharp claws, and smoke billowing from its mouth.",
|
||||
"Historical reimagining (image_1)": "Take the old photo of a historical figure in image_1 and reimagine them as a futuristic leader, wearing high-tech armor and standing in front of a modern city.",
|
||||
"Abstract art from a photograph (image_1)": "Turn the photograph in image_1 into an abstract work of art, where shapes and colors blend and distort, creating a completely new visual interpretation.",
|
||||
"Create an alien landscape (image_1)": "Use image_1, a barren desert, and transform it into an alien world with purple skies, strange plants, and glowing rocks scattered across the land."
|
||||
}
|
||||
}
|
||||
+1
-1
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
|
||||
|
||||
[project]
|
||||
name = "ComfyUI-OmniGen"
|
||||
version = "1.0.0"
|
||||
version = "1.2.0"
|
||||
description = "OmniGen node for AI image generation"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.8"
|
||||
|
||||
Reference in New Issue
Block a user