refactor: simplify to 2-param stdlib-only API
- Drop torch/numpy deps; node uses random/time/logging from stdlib only - Reduce surface to mode + seed inputs; remove sync_libraries, deterministic, overflow_behavior, use_torch_backend, batch_count - Sync version to 2.3.0 across pyproject.toml and __init__.py - Add pytest suite covering 4 modes + 64-bit wrap edges + IS_CHANGED contract - Rewrite README to match new API - Add uv.lock for reproducible dependency resolution
This commit is contained in:
@@ -4,185 +4,120 @@
|
||||
|
||||
[](https://opensource.org/licenses/Apache-2.0)
|
||||
[](https://github.com/comfyanonymous/ComfyUI)
|
||||
[](https://python.org)
|
||||
[](https://python.org)
|
||||
|
||||
**🎲 Advanced Seed Generator** - A professional-grade custom node for ComfyUI that provides comprehensive seed generation capabilities with multiple modes, state persistence, cross-library synchronization, and enterprise-level reliability features.
|
||||
**🎲 Advanced Seed Generator** — a small, focused custom node for ComfyUI that produces seed values in four modes: fixed, random, increment, and decrement. Pure Python standard library, no external dependencies.
|
||||
|
||||
## ✨ Features
|
||||
|
||||
- **🎯 Multiple Generation Modes**: Fixed, Random, Increment, and Decrement modes for various workflows
|
||||
- **🔄 State Persistence**: Maintains seed state across executions for increment/decrement modes
|
||||
- **🔗 Cross-Library Sync**: Synchronizes seeds across Python, NumPy, and PyTorch for consistent results
|
||||
- **⚡ Performance Optimized**: Intelligent backend selection (Python random vs PyTorch) based on batch size
|
||||
- **🛡️ Thread-Safe Operations**: Concurrent access protection with threading.RLock()
|
||||
- **🔧 Configurable Overflow**: Wrap, clamp, or error handling for boundary conditions
|
||||
- **📊 Batch Generation**: Generate up to 100,000 seeds efficiently in batch mode
|
||||
- **🎮 CUDA Support**: Full GPU acceleration support with deterministic mode options
|
||||
- **🐛 Comprehensive Logging**: Configurable debug logging for troubleshooting
|
||||
- **✅ Input Validation**: Robust error handling and parameter validation
|
||||
- **🎯 Four generation modes**: `fixed`, `random`, `increment`, `decrement`
|
||||
- **🔁 Cross-execution state**: increment/decrement remember the last seed across workflow runs
|
||||
- **🌐 Full 64-bit range**: seeds span `0` to `2⁶⁴ − 1` (`18,446,744,073,709,551,615`)
|
||||
- **♾️ Safe wrap-around**: increment past MAX wraps to `0`; decrement past `0` wraps to MAX
|
||||
- **📦 Zero dependencies**: uses only `random`, `time`, and `logging` from the Python stdlib
|
||||
- **🧩 Registry-ready**: published to the [ComfyUI Registry](https://registry.comfy.org)
|
||||
|
||||
## 🚀 Installation
|
||||
|
||||
### Method 1: Manual Installation
|
||||
1. Navigate to your ComfyUI `custom_nodes` directory
|
||||
2. Clone this repository:
|
||||
```bash
|
||||
git clone https://github.com/Limbicnation/ComfyUI-RandomSeedGenerator.git
|
||||
```
|
||||
3. Restart ComfyUI
|
||||
4. The node will appear under `utils` category as "🎲 Advanced Seed Generator"
|
||||
### Method 1 — Manual
|
||||
|
||||
### Method 2: ComfyUI Manager (Recommended)
|
||||
1. Install [ComfyUI Manager](https://github.com/ltdrdata/ComfyUI-Manager)
|
||||
2. Search for "Random Seed Generator" in the manager
|
||||
3. Install and restart ComfyUI
|
||||
```bash
|
||||
cd ComfyUI/custom_nodes
|
||||
git clone https://github.com/Limbicnation/ComfyUI-RandomSeedGenerator.git
|
||||
```
|
||||
|
||||
Restart ComfyUI. The node appears under **utils** as **🎲 Advanced Seed Generator**.
|
||||
|
||||
### Method 2 — ComfyUI Manager (recommended)
|
||||
|
||||
1. Install [ComfyUI Manager](https://github.com/ltdrdata/ComfyUI-Manager).
|
||||
2. Search for **Random Seed Generator**.
|
||||
3. Install and restart ComfyUI.
|
||||
|
||||
### Method 3 — ComfyUI Registry
|
||||
|
||||
### Method 3: ComfyUI Registry
|
||||
```bash
|
||||
comfy node install randomseedgenerator
|
||||
```
|
||||
|
||||
## 📖 Usage Guide
|
||||
## 📖 Usage
|
||||
|
||||
### Node Parameters
|
||||
### Parameters
|
||||
|
||||
| Parameter | Type | Default | Description |
|
||||
|-----------|------|---------|-------------|
|
||||
| **mode** | Dropdown | "fixed" | Generation mode: `fixed`, `increment`, `decrement`, `random` |
|
||||
| **seed** | Integer | 0 | Base seed value (0 to 18,446,744,073,709,551,615) |
|
||||
| **sync_libraries** | Boolean | True | Synchronize seed across Python, NumPy, PyTorch |
|
||||
| **deterministic** | Boolean | False | Enable full deterministic mode (may impact performance) |
|
||||
| **overflow_behavior** | Dropdown | "wrap" | Overflow handling: `wrap`, `clamp`, `error` |
|
||||
| **use_torch_backend** | Dropdown | "auto" | Backend selection: `auto`, `random`, `torch` |
|
||||
| **batch_count** | Integer | 1 | Number of seeds to generate (1-100,000) |
|
||||
| **mode** | Dropdown | `fixed` | One of `fixed`, `random`, `increment`, `decrement` |
|
||||
| **seed** | Integer | `0` | Used directly in `fixed` mode; ignored in the others |
|
||||
|
||||
### Generation Modes
|
||||
### Modes
|
||||
|
||||
#### 🔒 Fixed
|
||||
Returns the exact seed value provided. Use for reproducible generations.
|
||||
|
||||
#### 🔒 Fixed Mode
|
||||
Returns the exact seed value you specify. Perfect for reproducible generations.
|
||||
```
|
||||
Input: seed=12345 → Output: 12345 (always)
|
||||
seed=12345 → 12345 (always)
|
||||
```
|
||||
|
||||
#### 🎲 Random Mode
|
||||
Generates a new random seed on each execution (0 to 2^64-1).
|
||||
#### 🎲 Random
|
||||
Returns a fresh random seed in `[0, 2⁶⁴−1]` on every execution.
|
||||
|
||||
```
|
||||
Input: any seed → Output: 4831672946, 9573821047, ... (random)
|
||||
→ 4831672946…, 9573821047…, …
|
||||
```
|
||||
|
||||
#### ⬆️ Increment Mode
|
||||
Increments from the last generated seed by 1. State persists across workflow executions.
|
||||
#### ⬆️ Increment
|
||||
Returns `last_seed + 1`. State persists across workflow runs.
|
||||
|
||||
```
|
||||
First run: 42 → Second run: 43 → Third run: 44 ...
|
||||
run 1: 42 → run 2: 43 → run 3: 44 …
|
||||
```
|
||||
|
||||
#### ⬇️ Decrement Mode
|
||||
Decrements from the last generated seed by 1. State persists across workflow executions.
|
||||
#### ⬇️ Decrement
|
||||
Returns `last_seed − 1`. State persists across workflow runs.
|
||||
|
||||
```
|
||||
First run: 42 → Second run: 41 → Third run: 40 ...
|
||||
run 1: 42 → run 2: 41 → run 3: 40 …
|
||||
```
|
||||
|
||||
### Overflow Behavior Options
|
||||
### Wrap-around
|
||||
|
||||
- **🔄 Wrap (Default)**: Cycles around boundaries (MAX → MIN, MIN → MAX)
|
||||
- **🛑 Clamp**: Stops at boundaries (stays at MAX/MIN when limit reached)
|
||||
- **❌ Error**: Raises exception when overflow would occur
|
||||
Increment and decrement wrap at the 64-bit boundary: `MAX → 0` and `0 → MAX`. No configuration; this is always the behavior.
|
||||
|
||||
### Backend Selection
|
||||
### State scope
|
||||
|
||||
- **🤖 Auto (Recommended)**: Optimal backend selection based on batch size
|
||||
- Single seeds: Python random (fastest)
|
||||
- Batches ≥100: PyTorch CPU
|
||||
- Batches ≥1000: PyTorch GPU (if available)
|
||||
- **🐍 Random**: Force Python random module (good for small operations)
|
||||
- **🔥 Torch**: Force PyTorch backend (better for large batches)
|
||||
|
||||
## 💡 Usage Examples
|
||||
|
||||
### Basic Seed Generation
|
||||
1. Add "🎲 Advanced Seed Generator" to your workflow
|
||||
2. Set mode to "random" for exploration or "fixed" for reproducibility
|
||||
3. Connect the output to any node requiring a seed (KSampler, etc.)
|
||||
|
||||
### Batch Exploration Workflow
|
||||
1. Set `mode` to "increment"
|
||||
2. Set `batch_count` to 10
|
||||
3. Use with batch processors to generate variations systematically
|
||||
|
||||
### Professional Reproducibility Setup
|
||||
1. Set `mode` to "fixed"
|
||||
2. Enable `sync_libraries` and `deterministic`
|
||||
3. Document your seed values for exact reproduction
|
||||
|
||||
## 🔧 Advanced Configuration
|
||||
|
||||
### Environment Variables
|
||||
```bash
|
||||
# Set logging level for debugging
|
||||
export COMFYUI_SEED_LOG_LEVEL=DEBUG # Options: DEBUG, INFO, WARNING, ERROR
|
||||
```
|
||||
|
||||
### Performance Tuning
|
||||
- **Small batches (1-99)**: Use "random" backend for minimal overhead
|
||||
- **Medium batches (100-999)**: Use "auto" for optimal CPU performance
|
||||
- **Large batches (1000+)**: Use "auto" with CUDA for GPU acceleration
|
||||
The `_last_seed` counter is **class-level** — shared across every instance of the node in the current ComfyUI process. ComfyUI does not surface per-node identity to the executor, so isolated counters per node aren't possible without upstream API changes. If you need independent counters, use `random` mode and seed downstream samplers directly.
|
||||
|
||||
## 📋 Requirements
|
||||
|
||||
- **ComfyUI**: Latest version recommended
|
||||
- **Python**: 3.8 or higher
|
||||
- **Dependencies**:
|
||||
- `torch` (PyTorch)
|
||||
- `numpy`
|
||||
- `threading` (built-in)
|
||||
- `logging` (built-in)
|
||||
- **ComfyUI**: any current version
|
||||
- **Python**: 3.9+
|
||||
- **Dependencies**: none (Python stdlib only)
|
||||
|
||||
## 🔍 Troubleshooting
|
||||
|
||||
### Common Issues
|
||||
**Node not appearing in the menu**
|
||||
- Restart ComfyUI completely.
|
||||
- Check the ComfyUI console for import errors.
|
||||
|
||||
**Node not appearing in menu:**
|
||||
- Restart ComfyUI completely
|
||||
- Check console for import errors
|
||||
- Verify all dependencies are installed
|
||||
|
||||
**Increment/Decrement not working:**
|
||||
- State persists at class level - normal behavior
|
||||
- Use `reset_state()` method in console if needed
|
||||
- Check overflow_behavior setting
|
||||
|
||||
**Performance issues with large batches:**
|
||||
- Set backend to "torch" for batches >1000
|
||||
- Enable GPU if available for CUDA acceleration
|
||||
- Monitor memory usage with very large batches
|
||||
|
||||
### Debug Logging
|
||||
Enable detailed logging to diagnose issues:
|
||||
```bash
|
||||
export COMFYUI_SEED_LOG_LEVEL=DEBUG
|
||||
# Restart ComfyUI and check console output
|
||||
```
|
||||
**Increment/decrement counter is "wrong"**
|
||||
- The counter is shared across all instances of this node in the workflow. See *State scope* above.
|
||||
- To reset to `0`, call `AdvancedSeedGenerator.reset_state()` from the ComfyUI Python console.
|
||||
|
||||
## 🤝 Contributing
|
||||
|
||||
Contributions are welcome! Please:
|
||||
1. Fork the repository
|
||||
2. Create a feature branch
|
||||
3. Add tests for new functionality
|
||||
4. Submit a pull request
|
||||
Pull requests welcome. Please:
|
||||
|
||||
1. Fork and create a feature branch.
|
||||
2. Add tests under `tests/` for new behavior.
|
||||
3. Run `python -m pytest tests/`.
|
||||
4. Open a PR.
|
||||
|
||||
## 📝 License
|
||||
|
||||
This project is licensed under the Apache License 2.0 - see the [LICENSE](LICENSE) file for details.
|
||||
|
||||
## 🙏 Acknowledgments
|
||||
|
||||
- ComfyUI community for the amazing platform
|
||||
- Contributors and testers
|
||||
- Enhanced and maintained with Claude Code
|
||||
Apache License 2.0 — see [LICENSE](LICENSE).
|
||||
|
||||
---
|
||||
|
||||
**⭐ If this node helps your workflow, please consider starring the repository!**
|
||||
**⭐ If this node helps your workflow, please consider starring the repository.**
|
||||
|
||||
For issues, feature requests, or questions, please visit our [GitHub Issues](https://github.com/Limbicnation/ComfyUI-RandomSeedGenerator/issues) page.
|
||||
For issues or feature requests, see [GitHub Issues](https://github.com/Limbicnation/ComfyUI-RandomSeedGenerator/issues).
|
||||
|
||||
+6
-25
@@ -2,36 +2,17 @@
|
||||
Advanced Seed Generator Node for ComfyUI
|
||||
-----------------------------------------
|
||||
|
||||
An enhanced utility node that provides comprehensive seed generation
|
||||
for stable diffusion workflows. Supports multiple modes including
|
||||
fixed, random, increment, and decrement with robust error handling.
|
||||
|
||||
Features:
|
||||
- Multiple seed generation modes (fixed, random, increment, decrement)
|
||||
- Cross-library synchronization (Python, NumPy, PyTorch)
|
||||
- Comprehensive error handling and input validation
|
||||
- Configurable logging for debugging
|
||||
- Persistent state management across executions
|
||||
- CUDA support with deterministic mode options
|
||||
|
||||
Usage:
|
||||
1. Add the "Advanced Seed Generator" node to your workflow
|
||||
2. Connect it to nodes that require seed values (samplers, etc.)
|
||||
3. Choose your preferred generation mode
|
||||
4. Configure synchronization and deterministic options as needed
|
||||
Generates seed values for reproducible or exploratory image generation.
|
||||
Supports fixed, random, increment, and decrement modes.
|
||||
"""
|
||||
|
||||
__version__ = "2.2.0"
|
||||
__author__ = "Enhanced by Claude Code"
|
||||
__version__ = "2.3.0"
|
||||
__author__ = "Limbicnation"
|
||||
|
||||
from .random_seed_generator import AdvancedSeedGenerator, NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
|
||||
|
||||
# Re-export the mappings from the main module
|
||||
# This ensures consistency and single source of truth
|
||||
|
||||
# Module exports
|
||||
__all__ = [
|
||||
'NODE_CLASS_MAPPINGS',
|
||||
'NODE_CLASS_MAPPINGS',
|
||||
'NODE_DISPLAY_NAME_MAPPINGS',
|
||||
'AdvancedSeedGenerator'
|
||||
'AdvancedSeedGenerator',
|
||||
]
|
||||
+6
-5
@@ -1,12 +1,9 @@
|
||||
[project]
|
||||
name = "randomseedgenerator"
|
||||
description = "Advanced seed generator for ComfyUI with multiple modes, state persistence, and cross-library synchronization"
|
||||
version = "1.2.0"
|
||||
version = "2.3.0"
|
||||
license = {file = "LICENSE"}
|
||||
dependencies = [
|
||||
"torch>=1.9.0",
|
||||
"numpy>=1.19.0"
|
||||
]
|
||||
dependencies = []
|
||||
|
||||
[project.urls]
|
||||
Repository = "https://github.com/Limbicnation/ComfyUI-RandomSeedGenerator"
|
||||
@@ -17,3 +14,7 @@ PublisherId = "limbicnation"
|
||||
DisplayName = "ComfyUI-RandomSeedGenerator"
|
||||
Icon = "https://raw.githubusercontent.com/Limbicnation/ComfyUI-RandomSeedGenerator/main/icon.png"
|
||||
includes = []
|
||||
|
||||
[tool.pytest.ini_options]
|
||||
testpaths = ["tests"]
|
||||
addopts = "--confcutdir=tests --import-mode=importlib"
|
||||
|
||||
+61
-627
@@ -1,95 +1,32 @@
|
||||
"""ComfyUI node for reproducible seed generation with fixed, random, increment, and decrement modes."""
|
||||
|
||||
import random
|
||||
import time
|
||||
import logging
|
||||
import os
|
||||
import threading
|
||||
import numpy as np
|
||||
import torch
|
||||
from typing import Tuple, Union, Optional, Final, List
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Seed bounds (64-bit unsigned integer range)
|
||||
MIN_SEED = 0
|
||||
MAX_SEED = 0xFFFFFFFFFFFFFFFF # 2^64 - 1
|
||||
|
||||
|
||||
class AdvancedSeedGenerator:
|
||||
"""Generates seed values for reproducible or exploratory image generation.
|
||||
|
||||
Modes:
|
||||
fixed: Returns the exact seed value provided.
|
||||
random: Generates a new random seed (0 to 2^64-1) each execution.
|
||||
increment: Increments the last seed by 1 (wraps at MAX to 0).
|
||||
decrement: Decrements the last seed by 1 (wraps at 0 to MAX).
|
||||
|
||||
Note:
|
||||
_last_seed is class-level state shared across all node instances.
|
||||
ComfyUI does not expose per-node identity to IS_CHANGED or generate_seed,
|
||||
so per-instance isolation is not possible without upstream API changes.
|
||||
"""
|
||||
An advanced node that generates and synchronizes seed values for reproducible or exploratory image generation.
|
||||
|
||||
Features:
|
||||
- Multiple modes: fixed, increment, decrement, and random.
|
||||
- Robust state management for increment/decrement modes across executions.
|
||||
- Cross-library seed synchronization (random, numpy, torch).
|
||||
- Optional deterministic mode for complete reproducibility on CUDA devices.
|
||||
- Comprehensive error handling and input validation.
|
||||
- Configurable logging for debugging and monitoring.
|
||||
- Thread-safe operations for concurrent access.
|
||||
|
||||
Generation Modes:
|
||||
- fixed: Returns the exact seed value provided by user
|
||||
- random: Generates a new random seed (0 to 2^64-1) on each execution
|
||||
- increment: Increments the last generated seed by 1
|
||||
- decrement: Decrements the last generated seed by 1
|
||||
|
||||
Overflow Behavior (Configurable):
|
||||
Users can choose how increment/decrement operations handle boundary conditions:
|
||||
- \"wrap\" (default): Cycle around bounds (MAX -> MIN, MIN -> MAX)
|
||||
- \"clamp\": Stop at bounds (stay at MAX/MIN when limit reached)
|
||||
- \"error\": Raise ValueError exception when overflow would occur
|
||||
|
||||
This provides flexibility for different use cases while maintaining predictable behavior.
|
||||
|
||||
Cross-Library Compatibility:
|
||||
- Python random: Full 64-bit seed support
|
||||
- NumPy: Automatically truncates to 32-bit (logs when truncation occurs)
|
||||
- PyTorch: Full 64-bit seed support
|
||||
- CUDA: Full 64-bit seed support when available
|
||||
|
||||
Thread Safety:
|
||||
All state modifications are protected by threading.RLock() to ensure safe concurrent access.
|
||||
|
||||
Environment Variables:
|
||||
- COMFYUI_SEED_LOG_LEVEL: Set logging level (DEBUG, INFO, WARNING, ERROR)
|
||||
|
||||
Examples:
|
||||
>>> generator = AdvancedSeedGenerator()
|
||||
>>> result = generator.generate_seed("fixed", 42) # Returns (42,)
|
||||
>>> result = generator.generate_seed("increment", 0) # Returns (43,)
|
||||
>>> result = generator.generate_seed("random", 0) # Returns (random_value,)
|
||||
"""
|
||||
_last_seed = 0 # Class-level variable to store state across executions
|
||||
_logger = None # Class-level logger instance
|
||||
_lock = threading.RLock() # Thread-safe access to class state
|
||||
|
||||
# Constants for validation - Using Final for immutability
|
||||
MIN_SEED_VALUE: Final[int] = 0
|
||||
# 64-bit maximum (18,446,744,073,709,551,615) chosen for:
|
||||
# - Compatibility with modern diffusion models (Stable Diffusion, SDXL, etc.)
|
||||
# - Full range support for PyTorch generators
|
||||
# - Maximum entropy for random number generation
|
||||
# - Consistent with modern ML frameworks expecting 64-bit seeds
|
||||
MAX_SEED_VALUE: Final[int] = 0xffffffffffffffff
|
||||
DEFAULT_SEED: Final[int] = 0
|
||||
|
||||
# Configuration constants
|
||||
NUMPY_MAX_SEED: Final[int] = 2**32 - 1 # NumPy limited to 32-bit seeds (4,294,967,295)
|
||||
MAX_BATCH_COUNT: Final[int] = 100000 # Maximum number of seeds that can be generated in batch mode
|
||||
|
||||
# Backend selection thresholds for optimal performance
|
||||
TORCH_CUDA_BATCH_THRESHOLD: Final[int] = 1000 # Use GPU acceleration for batches >= 1000
|
||||
TORCH_CPU_BATCH_THRESHOLD: Final[int] = 100 # Use torch CPU for batches >= 100
|
||||
TORCH_BATCH_MIN_THRESHOLD: Final[int] = 10 # Minimum batch size to consider torch backend
|
||||
|
||||
@classmethod
|
||||
def _get_logger(cls):
|
||||
"""Get or create logger instance with configurable level."""
|
||||
if cls._logger is None:
|
||||
cls._logger = logging.getLogger(f"{__name__}.{cls.__name__}")
|
||||
log_level = os.environ.get('COMFYUI_SEED_LOG_LEVEL', 'WARNING')
|
||||
cls._logger.setLevel(getattr(logging, log_level.upper(), logging.WARNING))
|
||||
|
||||
if not cls._logger.handlers:
|
||||
handler = logging.StreamHandler()
|
||||
formatter = logging.Formatter('[%(name)s] %(levelname)s: %(message)s')
|
||||
handler.setFormatter(formatter)
|
||||
cls._logger.addHandler(handler)
|
||||
|
||||
return cls._logger
|
||||
_last_seed: int = 0
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
@@ -98,574 +35,71 @@ class AdvancedSeedGenerator:
|
||||
"mode": (["fixed", "increment", "decrement", "random"],),
|
||||
"seed": ("INT", {
|
||||
"default": 0,
|
||||
"min": 0,
|
||||
"max": 0xffffffffffffffff,
|
||||
"min": MIN_SEED,
|
||||
"max": MAX_SEED,
|
||||
"step": 1,
|
||||
"display": "number"
|
||||
}),
|
||||
"sync_libraries": ("BOOLEAN", {"default": True}),
|
||||
"deterministic": ("BOOLEAN", {"default": False}),
|
||||
"overflow_behavior": (["wrap", "clamp", "error"], {
|
||||
"default": "wrap",
|
||||
"tooltip": "How to handle overflow: wrap (cycle), clamp (stop at limits), error (raise exception)"
|
||||
}),
|
||||
"use_torch_backend": (["auto", "random", "torch"], {
|
||||
"default": "auto",
|
||||
"tooltip": "Random backend: auto (optimal), random (Python), torch (PyTorch)"
|
||||
}),
|
||||
"batch_count": ("INT", {
|
||||
"default": 1,
|
||||
"min": 1,
|
||||
"max": 100000, # Using literal value in INPUT_TYPES as class constants not accessible in classmethod
|
||||
"step": 1,
|
||||
"tooltip": "Number of seeds to generate (batch mode)"
|
||||
"display": "number",
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("INT", "INT")
|
||||
RETURN_NAMES = ("seed", "batch_count")
|
||||
RETURN_TYPES = ("INT",)
|
||||
RETURN_NAMES = ("seed",)
|
||||
FUNCTION = "generate_seed"
|
||||
CATEGORY = "utils"
|
||||
|
||||
def generate_seed(self, mode: str, seed: int, sync_libraries: bool = True, deterministic: bool = False, overflow_behavior: str = "wrap", use_torch_backend: str = "auto", batch_count: int = 1) -> Tuple[int, int]:
|
||||
"""
|
||||
Generate seed value(s) based on the selected mode and apply them if requested.
|
||||
def generate_seed(self, mode: str, seed: int) -> tuple[int]:
|
||||
"""Generate a seed value based on the selected mode.
|
||||
|
||||
Args:
|
||||
mode (str): The seed generation mode.
|
||||
seed (int): The user-defined seed (for 'fixed' mode).
|
||||
sync_libraries (bool): If True, synchronize the seed across Python, NumPy, and PyTorch.
|
||||
deterministic (bool): If True, enable full deterministic mode in PyTorch (may impact performance).
|
||||
overflow_behavior (str): How to handle overflow - "wrap", "clamp", or "error".
|
||||
use_torch_backend (str): Backend selection - "auto", "random", or "torch".
|
||||
batch_count (int): Number of seeds to generate (batch mode).
|
||||
mode: One of "fixed", "increment", "decrement", "random".
|
||||
seed: The user-provided seed value (used directly in fixed mode).
|
||||
|
||||
Returns:
|
||||
A tuple containing the generated integer seed and batch count.
|
||||
|
||||
Raises:
|
||||
ValueError: If mode is invalid, seed is out of bounds, or overflow occurs with "error" behavior.
|
||||
RuntimeError: If seed generation or library synchronization fails.
|
||||
"""
|
||||
logger = self._get_logger()
|
||||
|
||||
try:
|
||||
# Validate inputs
|
||||
self._validate_inputs(mode, seed, sync_libraries, deterministic, overflow_behavior, use_torch_backend, batch_count)
|
||||
|
||||
logger.debug(f"Generating seed with mode='{mode}', seed={seed}, sync={sync_libraries}, deterministic={deterministic}, backend='{use_torch_backend}', batch={batch_count}")
|
||||
|
||||
# Generate seed(s) based on mode and backend
|
||||
if batch_count == 1:
|
||||
final_seed = self._generate_seed_by_mode(mode, seed, overflow_behavior, use_torch_backend)
|
||||
else:
|
||||
# Batch mode - always returns the first seed for compatibility
|
||||
seeds = self._generate_seed_batch(mode, seed, batch_count, overflow_behavior, use_torch_backend)
|
||||
final_seed = seeds[0] if seeds else self.DEFAULT_SEED
|
||||
logger.info(f"Generated {len(seeds)} seeds in batch mode, using first seed: {final_seed}")
|
||||
|
||||
# Validate and clamp final seed
|
||||
final_seed = self._validate_and_clamp_seed(final_seed)
|
||||
|
||||
# Update class-level state (thread-safe)
|
||||
# Skip state update for batch increment/decrement as _generate_sequential_seed_batch already handled it
|
||||
if not (batch_count > 1 and mode in ['increment', 'decrement']):
|
||||
with self.__class__._lock:
|
||||
self.__class__._last_seed = final_seed
|
||||
if logger.isEnabledFor(logging.DEBUG):
|
||||
logger.debug(f"Updated _last_seed to {final_seed}")
|
||||
else:
|
||||
if logger.isEnabledFor(logging.DEBUG):
|
||||
logger.debug(f"Skipped _last_seed update for batch {mode} mode (already handled by batch generator)")
|
||||
|
||||
# Apply seed synchronization if requested
|
||||
if sync_libraries:
|
||||
self._apply_seed(final_seed, deterministic)
|
||||
if logger.isEnabledFor(logging.DEBUG):
|
||||
logger.debug(f"Applied seed {final_seed} to libraries")
|
||||
|
||||
if logger.isEnabledFor(logging.INFO):
|
||||
logger.info(f"Successfully generated seed: {final_seed} (mode: {mode})")
|
||||
return (final_seed, batch_count)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to generate seed: {str(e)}")
|
||||
# Return fallback seed to prevent complete failure
|
||||
fallback_seed = self.DEFAULT_SEED
|
||||
logger.warning(f"Using fallback seed: {fallback_seed}")
|
||||
return (fallback_seed, 1)
|
||||
|
||||
def _validate_inputs(self, mode: str, seed: int, sync_libraries: bool, deterministic: bool, overflow_behavior: str, use_torch_backend: str, batch_count: int) -> None:
|
||||
"""Validate all input parameters."""
|
||||
valid_modes = ["fixed", "increment", "decrement", "random"]
|
||||
valid_overflow_behaviors = ["wrap", "clamp", "error"]
|
||||
valid_backends = ["auto", "random", "torch"]
|
||||
|
||||
if not isinstance(mode, str) or mode not in valid_modes:
|
||||
raise ValueError(f"Invalid mode '{mode}'. Must be one of: {valid_modes}")
|
||||
|
||||
if not isinstance(seed, int):
|
||||
raise ValueError(f"Seed must be an integer, got {type(seed).__name__}")
|
||||
|
||||
if seed < self.MIN_SEED_VALUE or seed > self.MAX_SEED_VALUE:
|
||||
raise ValueError(f"Seed {seed} out of valid range [{self.MIN_SEED_VALUE}, {self.MAX_SEED_VALUE}]")
|
||||
|
||||
if not isinstance(sync_libraries, bool):
|
||||
raise ValueError(f"sync_libraries must be boolean, got {type(sync_libraries).__name__}")
|
||||
|
||||
if not isinstance(deterministic, bool):
|
||||
raise ValueError(f"deterministic must be boolean, got {type(deterministic).__name__}")
|
||||
|
||||
if not isinstance(overflow_behavior, str) or overflow_behavior not in valid_overflow_behaviors:
|
||||
raise ValueError(f"Invalid overflow_behavior '{overflow_behavior}'. Must be one of: {valid_overflow_behaviors}")
|
||||
|
||||
if not isinstance(use_torch_backend, str) or use_torch_backend not in valid_backends:
|
||||
raise ValueError(f"Invalid use_torch_backend '{use_torch_backend}'. Must be one of: {valid_backends}")
|
||||
|
||||
if not isinstance(batch_count, int) or batch_count < 1 or batch_count > self.MAX_BATCH_COUNT:
|
||||
raise ValueError(f"batch_count must be an integer between 1 and {self.MAX_BATCH_COUNT}, got {batch_count}")
|
||||
|
||||
def _generate_seed_by_mode(self, mode: str, seed: int, overflow_behavior: str = "wrap", use_torch_backend: str = "auto") -> int:
|
||||
"""
|
||||
Generate seed value based on the specified mode and backend.
|
||||
|
||||
Thread-safe generation with configurable overflow handling:
|
||||
- "wrap": Cycle around bounds (MAX -> MIN, MIN -> MAX)
|
||||
- "clamp": Stop at bounds (stay at MAX/MIN)
|
||||
- "error": Raise exception on overflow
|
||||
|
||||
Backend selection:
|
||||
- "auto": Use optimal backend (random for single seeds, torch for batches)
|
||||
- "random": Force Python random module
|
||||
- "torch": Force PyTorch backend
|
||||
"""
|
||||
try:
|
||||
if mode == 'fixed':
|
||||
return seed
|
||||
elif mode == 'random':
|
||||
return self._generate_random_seed(use_torch_backend)
|
||||
elif mode == 'increment':
|
||||
with self.__class__._lock:
|
||||
current_seed = self.__class__._last_seed
|
||||
new_seed = current_seed + 1
|
||||
return self._handle_overflow(new_seed, current_seed, "increment", overflow_behavior)
|
||||
elif mode == 'decrement':
|
||||
with self.__class__._lock:
|
||||
current_seed = self.__class__._last_seed
|
||||
new_seed = current_seed - 1
|
||||
return self._handle_overflow(new_seed, current_seed, "decrement", overflow_behavior)
|
||||
else:
|
||||
raise ValueError(f"Unsupported mode: {mode}")
|
||||
except Exception as e:
|
||||
raise RuntimeError(f"Failed to generate seed for mode '{mode}': {str(e)}")
|
||||
|
||||
def _generate_random_seed(self, use_torch_backend: str = "auto") -> int:
|
||||
"""
|
||||
Generate a single random seed using the specified backend.
|
||||
|
||||
Args:
|
||||
use_torch_backend (str): Backend preference - "auto", "random", or "torch"
|
||||
|
||||
Returns:
|
||||
int: Random seed value in valid range
|
||||
"""
|
||||
backend = self._select_optimal_backend(use_torch_backend, batch_size=1)
|
||||
logger = self._get_logger()
|
||||
|
||||
try:
|
||||
if backend == "torch":
|
||||
# Use torch.randint for direct integer generation (more stable)
|
||||
with torch.no_grad():
|
||||
# Use randint with safe range (PyTorch has limitations with very large ranges)
|
||||
# Use 48-bit range for better compatibility while maintaining good entropy
|
||||
torch_max = min(self.MAX_SEED_VALUE, 2**48 - 1)
|
||||
seed_tensor = torch.randint(
|
||||
low=self.MIN_SEED_VALUE,
|
||||
high=torch_max + 1,
|
||||
size=(1,),
|
||||
dtype=torch.int64,
|
||||
device='cpu'
|
||||
)
|
||||
return int(seed_tensor.item())
|
||||
else:
|
||||
# Use Python random (default, most efficient for single values)
|
||||
return random.randint(self.MIN_SEED_VALUE, self.MAX_SEED_VALUE)
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to generate random seed with {backend} backend: {str(e)}")
|
||||
# Fallback to Python random
|
||||
return random.randint(self.MIN_SEED_VALUE, self.MAX_SEED_VALUE)
|
||||
|
||||
def _generate_seed_batch(self, mode: str, seed: int, batch_count: int, overflow_behavior: str = "wrap", use_torch_backend: str = "auto") -> List[int]:
|
||||
"""
|
||||
Generate multiple seeds efficiently using batch operations.
|
||||
|
||||
Args:
|
||||
mode (str): The seed generation mode
|
||||
seed (int): Base seed value (for fixed mode)
|
||||
batch_count (int): Number of seeds to generate
|
||||
overflow_behavior (str): How to handle overflow
|
||||
use_torch_backend (str): Backend preference
|
||||
|
||||
Returns:
|
||||
List[int]: List of generated seed values
|
||||
"""
|
||||
logger = self._get_logger()
|
||||
|
||||
try:
|
||||
if mode == 'fixed':
|
||||
return [seed] * batch_count
|
||||
elif mode == 'random':
|
||||
return self._generate_random_seed_batch(batch_count, use_torch_backend)
|
||||
elif mode in ['increment', 'decrement']:
|
||||
return self._generate_sequential_seed_batch(mode, batch_count, overflow_behavior)
|
||||
else:
|
||||
raise ValueError(f"Unsupported mode for batch generation: {mode}")
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to generate seed batch: {str(e)}")
|
||||
# Fallback to single seed repeated
|
||||
fallback_seed = self._generate_seed_by_mode('fixed', self.DEFAULT_SEED, overflow_behavior, use_torch_backend)
|
||||
return [fallback_seed] * batch_count
|
||||
|
||||
def _generate_random_seed_batch(self, batch_count: int, use_torch_backend: str = "auto") -> List[int]:
|
||||
"""
|
||||
Generate multiple random seeds efficiently.
|
||||
|
||||
Args:
|
||||
batch_count (int): Number of seeds to generate
|
||||
use_torch_backend (str): Backend preference
|
||||
|
||||
Returns:
|
||||
List[int]: List of random seed values
|
||||
"""
|
||||
backend = self._select_optimal_backend(use_torch_backend, batch_size=batch_count)
|
||||
logger = self._get_logger()
|
||||
|
||||
try:
|
||||
if backend == "torch" and batch_count >= self.TORCH_BATCH_MIN_THRESHOLD:
|
||||
# Use torch for efficient batch generation
|
||||
with torch.no_grad():
|
||||
# Use CPU to avoid GPU memory overhead for small batches
|
||||
device = 'cpu'
|
||||
if batch_count >= self.TORCH_CUDA_BATCH_THRESHOLD and torch.cuda.is_available():
|
||||
device = 'cuda'
|
||||
|
||||
# Use randint with safe range for batch generation
|
||||
torch_max = min(self.MAX_SEED_VALUE, 2**48 - 1)
|
||||
seed_vals = torch.randint(
|
||||
low=self.MIN_SEED_VALUE,
|
||||
high=torch_max + 1,
|
||||
size=(batch_count,),
|
||||
dtype=torch.int64,
|
||||
device=device
|
||||
)
|
||||
|
||||
if device == 'cuda':
|
||||
seed_vals = seed_vals.cpu()
|
||||
|
||||
return seed_vals.tolist()
|
||||
else:
|
||||
# Use Python random for smaller batches or forced random backend
|
||||
return [random.randint(self.MIN_SEED_VALUE, self.MAX_SEED_VALUE) for _ in range(batch_count)]
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to generate batch with {backend} backend: {str(e)}")
|
||||
# Fallback to Python random
|
||||
return [random.randint(self.MIN_SEED_VALUE, self.MAX_SEED_VALUE) for _ in range(batch_count)]
|
||||
|
||||
def _generate_sequential_seed_batch(self, mode: str, batch_count: int, overflow_behavior: str) -> List[int]:
|
||||
"""
|
||||
Generate sequential seeds (increment/decrement) in batch.
|
||||
|
||||
Args:
|
||||
mode (str): "increment" or "decrement"
|
||||
batch_count (int): Number of seeds to generate
|
||||
overflow_behavior (str): How to handle overflow
|
||||
|
||||
Returns:
|
||||
List[int]: List of sequential seed values
|
||||
"""
|
||||
seeds = []
|
||||
|
||||
with self.__class__._lock:
|
||||
current_seed = self.__class__._last_seed
|
||||
|
||||
for i in range(batch_count):
|
||||
if mode == 'increment':
|
||||
new_seed = current_seed + 1
|
||||
else: # decrement
|
||||
new_seed = current_seed - 1
|
||||
|
||||
# Handle overflow for each step
|
||||
final_seed = self._handle_overflow(new_seed, current_seed, mode, overflow_behavior)
|
||||
seeds.append(final_seed)
|
||||
current_seed = final_seed
|
||||
|
||||
# Update the class state with the final seed
|
||||
self.__class__._last_seed = current_seed
|
||||
|
||||
return seeds
|
||||
|
||||
def _select_optimal_backend(self, use_torch_backend: str, batch_size: int = 1) -> str:
|
||||
"""
|
||||
Select the optimal backend based on preference and batch size.
|
||||
|
||||
Args:
|
||||
use_torch_backend (str): User preference - "auto", "random", or "torch"
|
||||
batch_size (int): Number of seeds to generate
|
||||
|
||||
Returns:
|
||||
str: Selected backend - "random" or "torch"
|
||||
"""
|
||||
if use_torch_backend == "random":
|
||||
return "random"
|
||||
elif use_torch_backend == "torch":
|
||||
return "torch"
|
||||
else: # "auto"
|
||||
# Auto-select based on batch size and PyTorch availability
|
||||
if batch_size >= self.TORCH_CUDA_BATCH_THRESHOLD and torch.cuda.is_available():
|
||||
return "torch"
|
||||
elif batch_size >= self.TORCH_CPU_BATCH_THRESHOLD: # Large batches benefit from torch even on CPU
|
||||
return "torch"
|
||||
else:
|
||||
return "random"
|
||||
|
||||
def _handle_overflow(self, new_seed: int, current_seed: int, operation: str, overflow_behavior: str) -> int:
|
||||
"""
|
||||
Handle overflow/underflow based on the specified behavior.
|
||||
|
||||
Args:
|
||||
new_seed (int): The calculated new seed value
|
||||
current_seed (int): The current seed value
|
||||
operation (str): "increment" or "decrement"
|
||||
overflow_behavior (str): "wrap", "clamp", or "error"
|
||||
|
||||
Returns:
|
||||
int: The final seed value after overflow handling
|
||||
|
||||
Raises:
|
||||
ValueError: If overflow_behavior is "error" and overflow occurs
|
||||
"""
|
||||
logger = self._get_logger()
|
||||
|
||||
# Check for overflow conditions
|
||||
if operation == "increment" and new_seed > self.MAX_SEED_VALUE:
|
||||
if overflow_behavior == "wrap":
|
||||
result = self.MIN_SEED_VALUE
|
||||
if logger.isEnabledFor(logging.DEBUG):
|
||||
logger.debug(f"Increment overflow: {current_seed} -> {result} (wrapped)")
|
||||
return result
|
||||
elif overflow_behavior == "clamp":
|
||||
result = self.MAX_SEED_VALUE
|
||||
if logger.isEnabledFor(logging.DEBUG):
|
||||
logger.debug(f"Increment overflow: {current_seed} -> {result} (clamped)")
|
||||
return result
|
||||
elif overflow_behavior == "error":
|
||||
raise ValueError(f"Increment overflow: seed {current_seed} + 1 exceeds maximum {self.MAX_SEED_VALUE}")
|
||||
|
||||
elif operation == "decrement" and new_seed < self.MIN_SEED_VALUE:
|
||||
if overflow_behavior == "wrap":
|
||||
result = self.MAX_SEED_VALUE
|
||||
if logger.isEnabledFor(logging.DEBUG):
|
||||
logger.debug(f"Decrement underflow: {current_seed} -> {result} (wrapped)")
|
||||
return result
|
||||
elif overflow_behavior == "clamp":
|
||||
result = self.MIN_SEED_VALUE
|
||||
if logger.isEnabledFor(logging.DEBUG):
|
||||
logger.debug(f"Decrement underflow: {current_seed} -> {result} (clamped)")
|
||||
return result
|
||||
elif overflow_behavior == "error":
|
||||
raise ValueError(f"Decrement underflow: seed {current_seed} - 1 is below minimum {self.MIN_SEED_VALUE}")
|
||||
|
||||
# No overflow occurred
|
||||
return new_seed
|
||||
|
||||
def _validate_and_clamp_seed(self, seed: int) -> int:
|
||||
"""Validate and clamp seed to valid range."""
|
||||
if not isinstance(seed, int):
|
||||
raise ValueError(f"Generated seed must be integer, got {type(seed).__name__}")
|
||||
|
||||
# Clamp to valid range
|
||||
clamped_seed = max(self.MIN_SEED_VALUE, min(seed, self.MAX_SEED_VALUE))
|
||||
|
||||
if clamped_seed != seed:
|
||||
self._get_logger().warning(f"Seed {seed} clamped to {clamped_seed}")
|
||||
|
||||
return clamped_seed
|
||||
A single-element tuple containing the generated seed.
|
||||
|
||||
def _apply_seed(self, seed_value: int, deterministic: bool = False) -> None:
|
||||
"""
|
||||
Apply the seed value across multiple libraries for consistent randomization.
|
||||
|
||||
Args:
|
||||
seed_value (int): The seed value to apply
|
||||
deterministic (bool): Whether to enable full deterministic mode
|
||||
|
||||
Raises:
|
||||
RuntimeError: If seed application fails for any library
|
||||
ValueError: If mode is not recognized.
|
||||
"""
|
||||
logger = self._get_logger()
|
||||
errors = []
|
||||
|
||||
# Apply seeds to all available libraries
|
||||
errors.extend(self._apply_python_seed(seed_value, logger))
|
||||
errors.extend(self._apply_numpy_seed(seed_value, logger))
|
||||
errors.extend(self._apply_pytorch_seed(seed_value, logger))
|
||||
errors.extend(self._apply_cuda_seeds(seed_value, deterministic, logger))
|
||||
|
||||
# Report any errors but don't fail completely
|
||||
if errors:
|
||||
logger.warning(f"Seed application completed with {len(errors)} errors: {'; '.join(errors)}")
|
||||
elif logger.isEnabledFor(logging.DEBUG):
|
||||
logger.debug("Successfully applied seed to all available libraries")
|
||||
|
||||
def _apply_python_seed(self, seed_value: int, logger: logging.Logger) -> list:
|
||||
"""Apply seed to Python's random module."""
|
||||
try:
|
||||
random.seed(seed_value)
|
||||
if logger.isEnabledFor(logging.DEBUG):
|
||||
logger.debug("Applied seed to Python random module")
|
||||
return []
|
||||
except Exception as e:
|
||||
error_msg = f"Failed to set Python random seed: {str(e)}"
|
||||
logger.error(error_msg)
|
||||
return [error_msg]
|
||||
|
||||
def _apply_numpy_seed(self, seed_value: int, logger: logging.Logger) -> list:
|
||||
"""Apply seed to NumPy's random module with 32-bit truncation."""
|
||||
try:
|
||||
# Handle potential overflow for NumPy (uses 32-bit seeds)
|
||||
numpy_seed = seed_value % self.NUMPY_MAX_SEED
|
||||
np.random.seed(numpy_seed)
|
||||
if numpy_seed != seed_value and logger.isEnabledFor(logging.DEBUG):
|
||||
logger.debug(f"NumPy seed truncated from {seed_value} to {numpy_seed}")
|
||||
if logger.isEnabledFor(logging.DEBUG):
|
||||
logger.debug("Applied seed to NumPy random")
|
||||
return []
|
||||
except Exception as e:
|
||||
error_msg = f"Failed to set NumPy random seed: {str(e)}"
|
||||
logger.error(error_msg)
|
||||
return [error_msg]
|
||||
|
||||
def _apply_pytorch_seed(self, seed_value: int, logger: logging.Logger) -> list:
|
||||
"""Apply seed to PyTorch."""
|
||||
try:
|
||||
torch.manual_seed(seed_value)
|
||||
if logger.isEnabledFor(logging.DEBUG):
|
||||
logger.debug("Applied seed to PyTorch")
|
||||
return []
|
||||
except Exception as e:
|
||||
error_msg = f"Failed to set PyTorch seed: {str(e)}"
|
||||
logger.error(error_msg)
|
||||
return [error_msg]
|
||||
|
||||
def _apply_cuda_seeds(self, seed_value: int, deterministic: bool, logger: logging.Logger) -> list:
|
||||
"""Apply CUDA seeds and configure deterministic mode."""
|
||||
errors = []
|
||||
|
||||
if not torch.cuda.is_available():
|
||||
if logger.isEnabledFor(logging.DEBUG):
|
||||
logger.debug("CUDA not available, skipping CUDA seed configuration")
|
||||
return errors
|
||||
|
||||
# Apply CUDA seed
|
||||
try:
|
||||
torch.cuda.manual_seed_all(seed_value)
|
||||
if logger.isEnabledFor(logging.DEBUG):
|
||||
logger.debug("Applied seed to CUDA")
|
||||
except Exception as e:
|
||||
error_msg = f"Failed to set CUDA seed: {str(e)}"
|
||||
logger.error(error_msg)
|
||||
errors.append(error_msg)
|
||||
|
||||
# Configure CUDNN deterministic mode
|
||||
try:
|
||||
torch.backends.cudnn.deterministic = deterministic
|
||||
torch.backends.cudnn.benchmark = not deterministic
|
||||
if logger.isEnabledFor(logging.DEBUG):
|
||||
logger.debug(f"Set CUDNN deterministic={deterministic}, benchmark={not deterministic}")
|
||||
except Exception as e:
|
||||
error_msg = f"Failed to configure CUDNN: {str(e)}"
|
||||
logger.error(error_msg)
|
||||
errors.append(error_msg)
|
||||
|
||||
return errors
|
||||
if mode == "fixed":
|
||||
result = seed
|
||||
elif mode == "random":
|
||||
result = random.randint(MIN_SEED, MAX_SEED)
|
||||
elif mode == "increment":
|
||||
result = self.__class__._last_seed + 1
|
||||
if result > MAX_SEED:
|
||||
result = MIN_SEED
|
||||
elif mode == "decrement":
|
||||
result = self.__class__._last_seed - 1
|
||||
if result < MIN_SEED:
|
||||
result = MAX_SEED
|
||||
else:
|
||||
raise ValueError(f"Unknown mode: {mode!r}")
|
||||
|
||||
self.__class__._last_seed = result
|
||||
logger.debug("Generated seed %d (mode=%s)", result, mode)
|
||||
return (result,)
|
||||
|
||||
@classmethod
|
||||
def IS_CHANGED(cls, mode: str, seed: int, sync_libraries: bool, deterministic: bool, overflow_behavior: str = "wrap", use_torch_backend: str = "auto", batch_count: int = 1) -> Union[float, str]:
|
||||
"""
|
||||
Force re-execution for modes that should produce a new result on each run.
|
||||
|
||||
Optimized for performance - minimal logging overhead.
|
||||
|
||||
Args:
|
||||
mode (str): The seed generation mode
|
||||
seed (int): The seed value (unused for dynamic modes)
|
||||
sync_libraries (bool): Whether libraries are synchronized
|
||||
deterministic (bool): Whether deterministic mode is enabled
|
||||
overflow_behavior (str): How to handle overflow (affects caching for increment/decrement)
|
||||
use_torch_backend (str): Backend preference for random generation
|
||||
batch_count (int): Number of seeds to generate
|
||||
|
||||
Returns:
|
||||
Union[float, str]: Unique value to force re-execution for dynamic modes,
|
||||
or constant for static modes
|
||||
"""
|
||||
try:
|
||||
# Dynamic modes always need re-execution
|
||||
if mode in ["random", "increment", "decrement"]:
|
||||
timestamp = time.time()
|
||||
# Only log if debug is explicitly enabled to avoid performance overhead
|
||||
logger = cls._get_logger()
|
||||
if logger.isEnabledFor(logging.DEBUG):
|
||||
logger.debug(f"IS_CHANGED: {timestamp} for dynamic mode '{mode}'")
|
||||
return timestamp
|
||||
else:
|
||||
# For fixed mode, return a stable cache key including all parameters
|
||||
cache_key = f"fixed_{seed}_{sync_libraries}_{deterministic}_{overflow_behavior}_{use_torch_backend}_{batch_count}"
|
||||
logger = cls._get_logger()
|
||||
if logger.isEnabledFor(logging.DEBUG):
|
||||
logger.debug(f"IS_CHANGED: '{cache_key}' for static mode '{mode}'")
|
||||
return cache_key
|
||||
except Exception as e:
|
||||
# Minimal error handling - don't call logger to avoid recursion
|
||||
print(f"[AdvancedSeedGenerator] Error in IS_CHANGED: {e}")
|
||||
# Fallback to timestamp to ensure execution
|
||||
def IS_CHANGED(cls, mode: str, seed: int):
|
||||
"""Force re-execution for dynamic modes; cache for fixed mode."""
|
||||
if mode in ("random", "increment", "decrement"):
|
||||
return time.time()
|
||||
|
||||
return f"fixed_{seed}"
|
||||
|
||||
@classmethod
|
||||
def reset_state(cls) -> None:
|
||||
"""Reset the class-level state. Useful for testing or reinitialization."""
|
||||
with cls._lock:
|
||||
cls._last_seed = cls.DEFAULT_SEED
|
||||
if cls._logger and cls._logger.isEnabledFor(logging.INFO):
|
||||
cls._logger.info("Reset AdvancedSeedGenerator state")
|
||||
|
||||
@classmethod
|
||||
def get_state_info(cls) -> dict:
|
||||
"""Get current state information for debugging."""
|
||||
with cls._lock:
|
||||
current_seed = cls._last_seed
|
||||
return {
|
||||
"last_seed": current_seed,
|
||||
"min_seed": cls.MIN_SEED_VALUE,
|
||||
"max_seed": cls.MAX_SEED_VALUE,
|
||||
"default_seed": cls.DEFAULT_SEED,
|
||||
"numpy_max_seed": cls.NUMPY_MAX_SEED,
|
||||
"logger_level": cls._logger.level if cls._logger else "Not initialized",
|
||||
"thread_safe": True
|
||||
}
|
||||
"""Reset class-level state. Useful for testing."""
|
||||
cls._last_seed = 0
|
||||
|
||||
|
||||
# ComfyUI Registration
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"AdvancedSeedGenerator": AdvancedSeedGenerator
|
||||
"AdvancedSeedGenerator": AdvancedSeedGenerator,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"AdvancedSeedGenerator": "🎲 Advanced Seed Generator"
|
||||
"AdvancedSeedGenerator": "🎲 Advanced Seed Generator",
|
||||
}
|
||||
|
||||
# Export for module-level access
|
||||
__all__ = ["AdvancedSeedGenerator", "NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
|
||||
+1
-2
@@ -1,2 +1 @@
|
||||
torch>=1.9.0
|
||||
numpy>=1.19.0
|
||||
# No external dependencies — uses only Python stdlib (random, time, logging)
|
||||
@@ -0,0 +1,128 @@
|
||||
"""Pytest suite for AdvancedSeedGenerator.
|
||||
|
||||
Covers the four modes (fixed, random, increment, decrement), wrap behavior at
|
||||
both 64-bit boundaries, IS_CHANGED contract, and reset_state.
|
||||
"""
|
||||
|
||||
import sys
|
||||
import time
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
# Allow `import random_seed_generator` when pytest is run from the repo root.
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
|
||||
|
||||
from random_seed_generator import ( # noqa: E402
|
||||
MAX_SEED,
|
||||
MIN_SEED,
|
||||
NODE_CLASS_MAPPINGS,
|
||||
NODE_DISPLAY_NAME_MAPPINGS,
|
||||
AdvancedSeedGenerator,
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _reset_class_state():
|
||||
"""`_last_seed` is class-level shared state — reset before every test."""
|
||||
AdvancedSeedGenerator.reset_state()
|
||||
yield
|
||||
AdvancedSeedGenerator.reset_state()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def gen():
|
||||
return AdvancedSeedGenerator()
|
||||
|
||||
|
||||
def test_fixed_returns_input_seed(gen):
|
||||
assert gen.generate_seed("fixed", 42) == (42,)
|
||||
|
||||
|
||||
def test_fixed_updates_last_seed(gen):
|
||||
gen.generate_seed("fixed", 12345)
|
||||
assert AdvancedSeedGenerator._last_seed == 12345
|
||||
|
||||
|
||||
def test_random_within_bounds(gen):
|
||||
(result,) = gen.generate_seed("random", 0)
|
||||
assert isinstance(result, int)
|
||||
assert MIN_SEED <= result <= MAX_SEED
|
||||
|
||||
|
||||
def test_random_produces_different_values(gen):
|
||||
samples = {gen.generate_seed("random", 0)[0] for _ in range(20)}
|
||||
# 20 draws from a 2^64 space colliding 19 times would be a miracle.
|
||||
assert len(samples) > 1
|
||||
|
||||
|
||||
def test_increment_advances_last_seed(gen):
|
||||
AdvancedSeedGenerator._last_seed = 100
|
||||
assert gen.generate_seed("increment", 0) == (101,)
|
||||
assert gen.generate_seed("increment", 0) == (102,)
|
||||
|
||||
|
||||
def test_decrement_decreases_last_seed(gen):
|
||||
AdvancedSeedGenerator._last_seed = 100
|
||||
assert gen.generate_seed("decrement", 0) == (99,)
|
||||
assert gen.generate_seed("decrement", 0) == (98,)
|
||||
|
||||
|
||||
def test_increment_wraps_at_max(gen):
|
||||
AdvancedSeedGenerator._last_seed = MAX_SEED
|
||||
assert gen.generate_seed("increment", 0) == (MIN_SEED,)
|
||||
|
||||
|
||||
def test_decrement_wraps_at_min(gen):
|
||||
AdvancedSeedGenerator._last_seed = MIN_SEED
|
||||
assert gen.generate_seed("decrement", 0) == (MAX_SEED,)
|
||||
|
||||
|
||||
def test_unknown_mode_raises_value_error(gen):
|
||||
with pytest.raises(ValueError, match="Unknown mode"):
|
||||
gen.generate_seed("teleport", 0)
|
||||
|
||||
|
||||
def test_is_changed_dynamic_modes_return_timestamp():
|
||||
for mode in ("random", "increment", "decrement"):
|
||||
before = time.time()
|
||||
result = AdvancedSeedGenerator.IS_CHANGED(mode, 0)
|
||||
after = time.time()
|
||||
assert isinstance(result, float)
|
||||
assert before <= result <= after
|
||||
|
||||
|
||||
def test_is_changed_fixed_returns_stable_key():
|
||||
assert AdvancedSeedGenerator.IS_CHANGED("fixed", 99) == "fixed_99"
|
||||
assert AdvancedSeedGenerator.IS_CHANGED("fixed", 99) == "fixed_99"
|
||||
|
||||
|
||||
def test_reset_state_clears_last_seed(gen):
|
||||
gen.generate_seed("fixed", 7777)
|
||||
assert AdvancedSeedGenerator._last_seed == 7777
|
||||
AdvancedSeedGenerator.reset_state()
|
||||
assert AdvancedSeedGenerator._last_seed == 0
|
||||
|
||||
|
||||
def test_node_registration_mappings():
|
||||
assert NODE_CLASS_MAPPINGS == {"AdvancedSeedGenerator": AdvancedSeedGenerator}
|
||||
assert "AdvancedSeedGenerator" in NODE_DISPLAY_NAME_MAPPINGS
|
||||
|
||||
|
||||
def test_input_types_schema_shape():
|
||||
schema = AdvancedSeedGenerator.INPUT_TYPES()
|
||||
assert "required" in schema
|
||||
assert set(schema["required"].keys()) == {"mode", "seed"}
|
||||
mode_choices, *_ = schema["required"]["mode"]
|
||||
assert set(mode_choices) == {"fixed", "random", "increment", "decrement"}
|
||||
seed_type, seed_meta = schema["required"]["seed"]
|
||||
assert seed_type == "INT"
|
||||
assert seed_meta["min"] == MIN_SEED
|
||||
assert seed_meta["max"] == MAX_SEED
|
||||
|
||||
|
||||
def test_comfyui_class_attributes():
|
||||
assert AdvancedSeedGenerator.RETURN_TYPES == ("INT",)
|
||||
assert AdvancedSeedGenerator.RETURN_NAMES == ("seed",)
|
||||
assert AdvancedSeedGenerator.FUNCTION == "generate_seed"
|
||||
assert AdvancedSeedGenerator.CATEGORY == "utils"
|
||||
Reference in New Issue
Block a user