diff --git a/UPDATE.md b/UPDATE.md
new file mode 100644
index 0000000..a9bac2b
--- /dev/null
+++ b/UPDATE.md
@@ -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
\ No newline at end of file
diff --git a/__init__.py b/__init__.py
index 68c2dba..41b23fe 100644
--- a/__init__.py
+++ b/__init__.py
@@ -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"
]
diff --git a/ailab_OmniGen.py b/ailab_OmniGen.py
index 878ce6e..00ecc28 100644
--- a/ailab_OmniGen.py
+++ b/ailab_OmniGen.py
@@ -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"
<|image_{i+1}|>" 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"
<|image_{i}|>"
- # 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"
<|image_{i+1}|>" 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"
<|image_{i}|>"
+ 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 🖼️"
+}
\ No newline at end of file
diff --git a/data.json b/data.json
new file mode 100644
index 0000000..0a805de
--- /dev/null
+++ b/data.json
@@ -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."
+ }
+}
\ No newline at end of file
diff --git a/pyproject.toml b/pyproject.toml
index 0b062b5..264fb77 100644
--- a/pyproject.toml
+++ b/pyproject.toml
@@ -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"