diff --git a/cli/civitai-search.py b/cli/civitai-search.py index 3d91e8711..cdb29c7ab 100755 --- a/cli/civitai-search.py +++ b/cli/civitai-search.py @@ -113,6 +113,11 @@ def search_civitai( return [] t0 = time.time() + import re + url_match = re.match(r'https?://civitai\.com/models/(\d+)', query.strip()) + if url_match: + query = url_match.group(1) + log.info(f'CivitAI: extracted model id={query} from URL') dct = { 'query': query } if len(tag) > 0: dct['tag'] = tag diff --git a/javascript/civitai.js b/javascript/civitai.js index d08fb6f28..ac8cae28c 100644 --- a/javascript/civitai.js +++ b/javascript/civitai.js @@ -11,6 +11,9 @@ String.prototype.format = function (args) { // eslint-disable-line no-extend-nat let selectedURL = ''; let selectedName = ''; let selectedType = ''; +let selectedBase = ''; +let selectedModelId = ''; +let selectedVersionId = ''; function clearModelDetails() { const el = gradioApp().getElementById('model-details') || gradioApp().getElementById('civitai_models_output') || gradioApp().getElementById('models_outcome'); @@ -84,7 +87,7 @@ async function modelCardClick(id) { data = data[0]; // assuming the first item is the one we want const versionsHTML = data.versions.map((v) => modelVersionsHTML.format({ - url: ``, + url: ``, name: v.name || 'unknown', type: v.files[0]?.type || 'unknown', base: v.base || 'unknown', @@ -113,11 +116,14 @@ async function modelCardClick(id) { el.innerHTML = modelHTML; } -function startCivitDownload(url, name, type) { - log('startCivitDownload', { url, name, type }); +function startCivitDownload(url, name, type, base, modelId, versionId) { + log('startCivitDownload', { url, name, type, base, modelId, versionId }); selectedURL = [url]; selectedName = [name]; selectedType = [type]; + selectedBase = [base || '']; + selectedModelId = [modelId || 0]; + selectedVersionId = [versionId || 0]; const civitDownloadBtn = gradioApp().getElementById('civitai_download_btn'); if (civitDownloadBtn) civitDownloadBtn.click(); } @@ -128,20 +134,44 @@ function startCivitAllDownload(evt) { selectedURL = []; selectedName = []; selectedType = []; + selectedBase = []; + selectedModelId = []; + selectedVersionId = []; for (const version of versions) { - const parsed = version.querySelector('td:nth-child(1) div')?.getAttribute('onclick')?.match(/startCivitDownload\('([^']+)', '([^']+)', '([^']+)'\)/); - if (!parsed || parsed.length < 4) continue; + const parsed = version.querySelector('td:nth-child(1) div')?.getAttribute('onclick')?.match(/startCivitDownload\('([^']+)', '([^']+)', '([^']+)', '([^']*)', (\d+), (\d+)\)/); + if (!parsed || parsed.length < 7) continue; selectedURL.push(parsed[1]); selectedName.push(parsed[2]); selectedType.push(parsed[3]); + selectedBase.push(parsed[4]); + selectedModelId.push(parseInt(parsed[5], 10)); + selectedVersionId.push(parseInt(parsed[6], 10)); } const civitDownloadBtn = gradioApp().getElementById('civitai_download_btn'); if (civitDownloadBtn) civitDownloadBtn.click(); } -function downloadCivitModel(modelUrl, modelName, modelType, modelPath, civitToken, innerHTML) { - log('downloadCivitModel', { modelUrl, modelName, modelType, modelPath, civitToken }); +function downloadCivitModel(modelUrl, modelName, modelType, modelBase, mId, vId, modelPath, civitToken, innerHTML) { + log('downloadCivitModel', { modelUrl, modelName, modelType, modelBase, mId, vId, modelPath, civitToken }); const el = gradioApp().getElementById('civitai_models_output') || gradioApp().getElementById('models_outcome'); const currentHTML = el?.innerHTML || ''; - return [selectedURL, selectedName, selectedType, modelPath, civitToken, currentHTML]; + return [selectedURL, selectedName, selectedType, selectedBase, selectedModelId, selectedVersionId, modelPath, civitToken, currentHTML]; } + +let civitMutualExcludeBound = false; + +function civitaiMutualExclude() { + if (civitMutualExcludeBound) return; + const searchEl = gradioApp().querySelector('#civit_search_text textarea'); + const tagEl = gradioApp().querySelector('#civit_search_tag textarea'); + if (!searchEl || !tagEl) return; + civitMutualExcludeBound = true; + searchEl.addEventListener('input', () => { + tagEl.closest('.gradio-textbox')?.classList.toggle('disabled-look', !!searchEl.value.trim()); + }); + tagEl.addEventListener('input', () => { + searchEl.closest('.gradio-textbox')?.classList.toggle('disabled-look', !!tagEl.value.trim()); + }); +} + +onUiLoaded(civitaiMutualExclude); diff --git a/javascript/sdnext.css b/javascript/sdnext.css index fea7396ac..c34b5c951 100644 --- a/javascript/sdnext.css +++ b/javascript/sdnext.css @@ -2236,6 +2236,11 @@ div:has(>#tab-gallery-folders) { filter: blur(0); } +.disabled-look textarea { + opacity: 0.4; + pointer-events: none; +} + @keyframes spin { from { transform: rotate(0deg); diff --git a/modules/civitai/download_civitai.py b/modules/civitai/download_civitai.py index b2e809d06..da5378b41 100644 --- a/modules/civitai/download_civitai.py +++ b/modules/civitai/download_civitai.py @@ -406,14 +406,24 @@ def download_civit_preview(model_path: str, preview_url: str): return 200, str(total_size), '' -def download_civit_model(model_url: str, model_name: str = '', model_path: str = '', model_type: str = '', token: str = None): +def download_civit_model(model_url: str, model_name: str = '', model_path: str = '', model_type: str = '', token: str = None, + base_model: str = '', model_id: int = 0, version_id: int = 0): """Legacy function — delegates to DownloadManager for non-blocking downloads.""" if not model_url: log.error('Model download: no url provided') return None + if not version_id: + import re + match = re.search(r'/api/download/models/(\d+)', model_url) + if match: + version_id = int(match.group(1)) from modules.civitai.filemanage_civitai import get_type_folder if not model_path: - folder = str(get_type_folder(model_type or 'Checkpoint')) + if getattr(shared.opts, 'civitai_save_subfolder_enabled', False): + from modules.civitai.filemanage_civitai import resolve_save_path + folder = str(resolve_save_path(model_type or 'Checkpoint', model_name=model_name, base_model=base_model)) + else: + folder = str(get_type_folder(model_type or 'Checkpoint')) elif os.path.isabs(model_path): folder = model_path else: @@ -424,6 +434,8 @@ def download_civit_model(model_url: str, model_name: str = '', model_path: str = filename=model_name or "Unknown", model_type=model_type, token=token, + model_id=model_id, + version_id=version_id, ) # Wait for completion (legacy blocking behavior) while item.status in ("queued", "downloading", "verifying"): diff --git a/modules/civitai/search_civitai.py b/modules/civitai/search_civitai.py index 07fc7f809..e3b116dc2 100644 --- a/modules/civitai/search_civitai.py +++ b/modules/civitai/search_civitai.py @@ -1,3 +1,4 @@ +import re import time from installer import log from modules.civitai.client_civitai import client @@ -20,12 +21,18 @@ def search_civitai( token: str = None, exact: bool = True, ) -> list[CivitModel]: - if not query: - log.error('CivitAI: empty query') + if not query and not tag and not sort: + log.error('CivitAI: no search criteria provided') return [] t0 = time.time() + # URL query → extract model ID (e.g. https://civitai.com/models/967405/nova-orange-xl) + url_match = re.match(r'https?://civitai\.com/models/(\d+)', query.strip()) + if url_match: + query = url_match.group(1) + log.info(f'CivitAI: extracted model id={query} from URL') + # Numeric query → single model fetch if query.isnumeric(): model = client.get_model(int(query), token=token) @@ -49,7 +56,7 @@ def search_civitai( all_models = response.items exact_models: list[CivitModel] = [] - if exact: + if exact and query: q_lower = query.lower() for model in all_models: names = [model.name.lower()] diff --git a/modules/ui_models.py b/modules/ui_models.py index 34b9b9798..2856448e6 100644 --- a/modules/ui_models.py +++ b/modules/ui_models.py @@ -483,11 +483,21 @@ def create_ui(): outputs=[models_outcome] ) - with gr.Tab(label="CivitAI", elem_id="models_civitai_tab"): + with gr.Tab(label="CivitAI", elem_id="models_civitai_tab") as civitai_tab: from modules.civitai.search_civitai import search_civitai, create_model_cards, base_models - def civitai_search(civit_search_text, civit_search_tag, civit_nsfw, civit_type, civit_base, civit_token): - results = search_civitai(query=civit_search_text, tag=civit_search_tag, nsfw=civit_nsfw, types=civit_type, base=civit_base, token=civit_token) + sort_fallback = ['', 'Most Downloaded', 'Highest Rated', 'Most Liked', 'Most Discussed', + 'Most Collected', 'Most Images', 'Newest', 'Oldest'] + type_fallback = ['', 'Checkpoint', 'TextualInversion', 'Hypernetwork', 'AestheticGradient', + 'LORA', 'LoCon', 'DoRA', 'Controlnet', 'Upscaler', 'MotionModule', + 'VAE', 'Poses', 'Wildcards', 'Workflows', 'Detection', 'Other'] + def civitai_search(civit_search_text, civit_search_tag, civit_nsfw, civit_type, + civit_base, civit_token, civit_sort, civit_period): + if civit_search_text and civit_search_tag: + civit_search_tag = '' # query+tag is broken at API level, keyword wins + results = search_civitai(query=civit_search_text, tag=civit_search_tag, nsfw=civit_nsfw, + types=civit_type, base=civit_base, token=civit_token, + sort=civit_sort, period=civit_period) html = create_model_cards(results) return html @@ -496,50 +506,115 @@ def create_ui(): opts.civitai_token = token opts.save() - def civitai_download(model_urls, model_names, model_types, model_path, civit_token, model_output): + def civitai_download(model_urls, model_names, model_types, model_bases, + model_ids, version_ids, model_path, civit_token, model_output): from modules.civitai.download_civitai import download_civit_model - for model_url, model_name, model_type in zip(model_urls, model_names, model_types, strict=False): + for model_url, model_name, model_type, model_base, model_id, version_id in zip( + model_urls, model_names, model_types, model_bases, model_ids, version_ids, strict=False): msg = f"

Initiating download

{model_name} | {model_type} | {model_url}

" yield msg + model_output - download_civit_model(model_url, model_name, model_path, model_type, civit_token) + download_civit_model(model_url, model_name, model_path, model_type, civit_token, + base_model=model_base, model_id=int(model_id or 0), version_id=int(version_id or 0)) yield model_output + def civitai_toggle_subfolder(enabled, template): + opts.data['civitai_save_subfolder_enabled'] = enabled + if enabled and not template: + template = '{{BASEMODEL}}' + opts.data['civitai_save_subfolder'] = template + opts.save() + return gr.update(value=template, interactive=enabled) + with gr.Row(): gr.HTML('

Search & Download

') with gr.Row(elem_id='civitai_search_row'): - civit_search_text = gr.Textbox(label='', placeholder='keyword', elem_id="civit_search_text") - civit_search_tag = gr.Textbox(label='', placeholder='tag', elem_id="civit_search_text") + civit_search_text = gr.Textbox(label='', placeholder='keyword, model id, or civitai url', elem_id="civit_search_text") + civit_search_tag = gr.Textbox(label='', placeholder='tag', elem_id="civit_search_tag") civit_search_text_btn = ToolButton(value=ui_symbols.search, interactive=True, elem_id="civit_text_search") - with gr.Accordion(label='Advanced', open=False, elem_id="civitai_search_options"): + with gr.Accordion(label='Options', open=False, elem_id="civitai_search_options"): civit_download_btn = gr.Button(value="Download model", variant='primary', elem_id="civitai_download_btn", visible=False) with gr.Row(): - civit_token = gr.Textbox(opts.civitai_token, label='CivitAI token', placeholder='optional access token for private or gated models', elem_id="civitai_token") + civit_type = gr.Dropdown(choices=type_fallback, label='Model type', value='', elem_id='civit_type') + civit_base = gr.Dropdown(choices=base_models, label='Base model', value='') + with gr.Row(): + civit_sort = gr.Dropdown(choices=sort_fallback, label='Sort', value='', elem_id='civit_sort') + civit_period = gr.Dropdown( + choices=['', 'AllTime', 'Year', 'Month', 'Week', 'Day'], + label='Time period', value='', elem_id='civit_period', + ) with gr.Row(): civit_nsfw = gr.Checkbox(label='NSFW allowed', value=True) with gr.Row(): - civit_type = gr.Textbox(label='Target model type', placeholder='Checkpoint, LORA, ...', value='') - with gr.Row(): - # civit_base = gr.Textbox(label='Base model', placeholder='SDXL, ...') - civit_base = gr.Dropdown(choices=base_models, label='Base model', value='') + civit_token = gr.Textbox(opts.civitai_token, label='CivitAI token', placeholder='optional access token for private or gated models', elem_id="civitai_token") with gr.Row(): civit_folder = gr.Textbox(label='Download folder', placeholder='optional folder for downloads') + with gr.Row(): + civit_subfolder_enabled = gr.Checkbox( + label='Sort downloads into subfolders', + value=getattr(opts, 'civitai_save_subfolder_enabled', False), + elem_id='civit_subfolder_enabled', + ) + civit_subfolder_template = gr.Textbox( + label='Subfolder template', + value=getattr(opts, 'civitai_save_subfolder', '') if getattr(opts, 'civitai_save_subfolder_enabled', False) else '', + placeholder='e.g. {{BASEMODEL}} or {{CREATOR}}/{{BASEMODEL}}', + interactive=getattr(opts, 'civitai_save_subfolder_enabled', False), + elem_id='civit_subfolder_template', + ) with gr.Row(): civitai_models_output = gr.HTML('', elem_id="civitai_models_output") - # sort, period, limit - _dummy = gr.Label(visible=False) # dummy component to get argspec later - civit_inputs = [civit_search_text, civit_search_tag, civit_nsfw, civit_type, civit_base, civit_token] + _dummy = gr.Label(visible=False) + civit_inputs = [civit_search_text, civit_search_tag, civit_nsfw, civit_type, + civit_base, civit_token, civit_sort, civit_period] civit_search_text_btn.click(fn=civitai_search, inputs=civit_inputs, outputs=[civitai_models_output]) civit_search_text.submit(fn=civitai_search, inputs=civit_inputs, outputs=[civitai_models_output]) civit_search_tag.submit(fn=civitai_search, inputs=civit_inputs, outputs=[civitai_models_output]) civit_token.change(fn=civitai_update_token, inputs=[civit_token], outputs=[]) + civit_subfolder_enabled.change( + fn=civitai_toggle_subfolder, + inputs=[civit_subfolder_enabled, civit_subfolder_template], + outputs=[civit_subfolder_template], + ) + civit_subfolder_template.change( + fn=lambda v: (setattr(opts, 'civitai_save_subfolder', v), opts.save()), + inputs=[civit_subfolder_template], outputs=[], + ) civit_download_btn.click( fn=civitai_download, _js="downloadCivitModel", - inputs=[_dummy, _dummy, _dummy, civit_folder, civit_token, civitai_models_output], + inputs=[_dummy, _dummy, _dummy, _dummy, _dummy, _dummy, civit_folder, civit_token, civitai_models_output], outputs=[civitai_models_output], show_progress='full', ) + _civitai_loaded = False + + def civitai_on_tab_enter(): + nonlocal _civitai_loaded + if _civitai_loaded: + return [gr.update(), gr.update(), gr.update(), gr.update(), gr.update()] + _civitai_loaded = True + from modules.civitai.client_civitai import client + options = client.discover_options() + type_choices = [''] + (options.get('types', []) or type_fallback[1:]) + sort_choices = [''] + (options.get('sort', []) or sort_fallback[1:]) + base_choices = [''] + (options.get('base_models', []) or base_models[1:]) + results = search_civitai(query='', sort='Most Downloaded', period='AllTime', limit=20) + html = create_model_cards(results) + return [ + gr.update(choices=type_choices), + gr.update(choices=sort_choices, value='Most Downloaded'), + gr.update(choices=base_choices), + gr.update(value='AllTime'), + html, + ] + + civitai_tab.select( + fn=civitai_on_tab_enter, + inputs=[], + outputs=[civit_type, civit_sort, civit_base, civit_period, civitai_models_output], + ) + with gr.Tab(label="Huggingface", elem_id="models_huggingface_tab"): from modules.models_hf import hf_search, hf_select, hf_download_model, hf_update_token with gr.Column(scale=6): diff --git a/test/test-detailer-api.py b/test/test-detailer-api.py new file mode 100644 index 000000000..acc4622d7 --- /dev/null +++ b/test/test-detailer-api.py @@ -0,0 +1,653 @@ +#!/usr/bin/env python +""" +API tests for YOLO Detailer endpoints. + +Tests: +- GET /sdapi/v1/detailers — model enumeration +- POST /sdapi/v1/detect — object detection on test images +- POST /sdapi/v1/txt2img — generation with detailer enabled + +Requires a running SD.Next instance with a model loaded. + +Usage: + python test/test-detailer-api.py [--url URL] [--image PATH] +""" + +import io +import os +import sys +import time +import json +import base64 +import argparse +import requests +import urllib3 + +urllib3.disable_warnings(urllib3.exceptions.InsecureRequestWarning) + +# Reference model cover images with faces (best for detailer testing) +FACE_TEST_IMAGES = [ + 'models/Reference/ponyRealism_V23.jpg', # realistic woman, clear face + 'models/Reference/HiDream-ai--HiDream-I1-Fast.jpg', # realistic man, clear face + text + 'models/Reference/stabilityai--stable-diffusion-xl-base-1.0.jpg', # realistic woman portrait + 'models/Reference/CalamitousFelicitousness--Anima-sdnext-diffusers.jpg', # anime face (non-realistic test) +] + +# Fallback images (no guaranteed faces) +FALLBACK_IMAGES = [ + 'html/sdnext-robot-2k.jpg', + 'html/favicon.png', +] + + +class DetailerAPITest: + """Test harness for YOLO Detailer API endpoints.""" + + def __init__(self, base_url, image_path=None, timeout=300): + self.base_url = base_url.rstrip('/') + self.test_images = {} # name -> base64 + self.timeout = timeout + self.results = { + 'enumerate': {'passed': 0, 'failed': 0, 'skipped': 0, 'tests': []}, + 'detect': {'passed': 0, 'failed': 0, 'skipped': 0, 'tests': []}, + 'generate': {'passed': 0, 'failed': 0, 'skipped': 0, 'tests': []}, + 'detailer_params': {'passed': 0, 'failed': 0, 'skipped': 0, 'tests': []}, + } + self._category = 'enumerate' + self._critical_error = None + self._load_images(image_path) + + def _encode_image(self, path): + from PIL import Image + image = Image.open(path) + if image.mode == 'RGBA': + image = image.convert('RGB') + buf = io.BytesIO() + image.save(buf, 'JPEG') + return base64.b64encode(buf.getvalue()).decode(), image.size + + def _load_images(self, image_path=None): + if image_path and os.path.exists(image_path): + b64, size = self._encode_image(image_path) + name = os.path.basename(image_path) + self.test_images[name] = b64 + print(f" Test image: {image_path} ({size})") + return + + # Load all available face test images + for p in FACE_TEST_IMAGES: + if os.path.exists(p): + b64, size = self._encode_image(p) + name = os.path.basename(p) + self.test_images[name] = b64 + print(f" Loaded: {name} ({size[0]}x{size[1]})") + + # Fallback if no face images found + if not self.test_images: + for p in FALLBACK_IMAGES: + if os.path.exists(p): + b64, size = self._encode_image(p) + name = os.path.basename(p) + self.test_images[name] = b64 + print(f" Fallback: {name} ({size[0]}x{size[1]})") + break + + if not self.test_images: + print(" WARNING: No test images found, detect tests will be skipped") + + @property + def image_b64(self): + """Return the first available test image for backwards compat.""" + if self.test_images: + return next(iter(self.test_images.values())) + return None + + def _get(self, endpoint): + try: + r = requests.get(f'{self.base_url}{endpoint}', timeout=self.timeout, verify=False) + if r.status_code != 200: + return {'error': r.status_code, 'reason': r.reason} + return r.json() + except requests.exceptions.ConnectionError: + return {'error': 'connection_refused', 'reason': 'Server not running'} + except Exception as e: + return {'error': 'exception', 'reason': str(e)} + + def _post(self, endpoint, data): + try: + r = requests.post(f'{self.base_url}{endpoint}', json=data, timeout=self.timeout, verify=False) + if r.status_code != 200: + return {'error': r.status_code, 'reason': r.reason} + return r.json() + except requests.exceptions.ConnectionError: + return {'error': 'connection_refused', 'reason': 'Server not running'} + except Exception as e: + return {'error': 'exception', 'reason': str(e)} + + def record(self, passed, name, detail=''): + status = 'PASS' if passed else 'FAIL' + self.results[self._category]['passed' if passed else 'failed'] += 1 + self.results[self._category]['tests'].append((status, name)) + msg = f' {status}: {name}' + if detail: + msg += f' ({detail})' + print(msg) + + def skip(self, name, reason): + self.results[self._category]['skipped'] += 1 + self.results[self._category]['tests'].append(('SKIP', name)) + print(f' SKIP: {name} ({reason})') + + # ========================================================================= + # Tests: Model Enumeration + # ========================================================================= + + def test_detailers_list(self): + """GET /sdapi/v1/detailers returns a list of available models.""" + self._category = 'enumerate' + print("\n--- Detailer Model Enumeration ---") + + data = self._get('/sdapi/v1/detailers') + if 'error' in data: + self.record(False, 'detailers_list', f"error: {data}") + self._critical_error = f"Server error: {data}" + return [] + + if not isinstance(data, list): + self.record(False, 'detailers_list', f"expected list, got {type(data).__name__}") + return [] + + self.record(True, 'detailers_list', f"{len(data)} models found") + + # Verify each entry has expected fields + if len(data) > 0: + sample = data[0] + has_name = 'name' in sample + self.record(has_name, 'detailer_entry_has_name', f"sample: {sample}") + if not has_name: + self.record(False, 'detailer_entry_schema', "missing 'name' field") + + return data + + # ========================================================================= + # Tests: Detection + # ========================================================================= + + def _validate_detect_response(self, data, label): + """Validate detection response schema and return detection count.""" + expected_keys = ['classes', 'labels', 'boxes', 'scores'] + for key in expected_keys: + if key not in data: + self.record(False, f'{label}_schema_{key}', f"missing '{key}'") + return -1 + + # All arrays should have the same length + lengths = [len(data[key]) for key in expected_keys] + all_same = len(set(lengths)) <= 1 + if not all_same: + self.record(False, f'{label}_array_lengths', f"mismatched: {dict(zip(expected_keys, lengths))}") + return -1 + + n = lengths[0] + + if n > 0: + # Scores should be in [0, 1] + scores_valid = all(0 <= s <= 1 for s in data['scores']) + if not scores_valid: + self.record(False, f'{label}_scores_range', f"scores: {data['scores']}") + + # Boxes should be lists of 4 numbers + boxes_valid = all(isinstance(b, list) and len(b) == 4 for b in data['boxes']) + if not boxes_valid: + self.record(False, f'{label}_boxes_format', "bad box format") + + return n + + # Face detection models to try (in priority order) + FACE_MODELS = ['face-yolo8n', 'face-yolo8m', 'anzhc-face-1024-seg-8n'] + + def _pick_face_model(self, available_models): + """Pick the best face detection model from available ones.""" + available_names = [m.get('name', '') for m in available_models] if available_models else [] + for model in self.FACE_MODELS: + if model in available_names: + return model + return '' # fall back to server default + + def test_detect_all_images(self, available_models=None): + """POST /sdapi/v1/detect on each loaded test image with a face model.""" + self._category = 'detect' + print("\n--- Detection Tests (per-image) ---") + + if not self.test_images: + self.skip('detect_all', 'no test images') + return + + if self._critical_error: + self.skip('detect_all', self._critical_error) + return + + face_model = self._pick_face_model(available_models) + if face_model: + print(f" Using face model: {face_model}") + else: + print(" No face model available, using server default") + + total_detections = 0 + any_face_found = False + + for img_name, img_b64 in self.test_images.items(): + short = img_name.replace('.jpg', '')[:40] + data = self._post('/sdapi/v1/detect', {'image': img_b64, 'model': face_model}) + + if 'error' in data: + self.record(False, f'detect_{short}', f"error: {data}") + continue + + n = self._validate_detect_response(data, f'detect_{short}') + if n < 0: + continue + + labels = data.get('labels', []) + scores = data.get('scores', []) + detail_parts = [f"{n} detections"] + if labels: + detail_parts.append(f"labels={labels}") + if scores: + detail_parts.append(f"top_score={max(scores):.3f}") + + self.record(True, f'detect_{short}', ', '.join(detail_parts)) + total_detections += n + if n > 0: + any_face_found = True + + self.record(any_face_found, 'detect_found_faces', + f"{total_detections} total detections across {len(self.test_images)} images") + + def test_detect_with_model(self, model_name): + """POST /sdapi/v1/detect with a specific model on all images.""" + if not self.test_images: + self.skip(f'detect_model_{model_name}', 'no test images') + return + + total = 0 + for _img_name, img_b64 in self.test_images.items(): + data = self._post('/sdapi/v1/detect', {'image': img_b64, 'model': model_name}) + if 'error' not in data: + total += len(data.get('scores', [])) + + self.record(True, f'detect_model_{model_name}', f"{total} detections across {len(self.test_images)} images") + + # ========================================================================= + # Tests: Generation with Detailer + # ========================================================================= + + def test_txt2img_with_detailer(self): + """POST /sdapi/v1/txt2img with detailer_enabled=True.""" + self._category = 'generate' + print("\n--- Generation with Detailer ---") + + if self._critical_error: + self.skip('txt2img_detailer', self._critical_error) + return + + payload = { + 'prompt': 'a photo of a person, face, portrait', + 'negative_prompt': '', + 'steps': 10, + 'width': 512, + 'height': 512, + 'seed': 42, + 'save_images': False, + 'send_images': True, + 'detailer_enabled': True, + 'detailer_strength': 0.3, + 'detailer_steps': 5, + 'detailer_conf': 0.3, + 'detailer_max': 3, + } + + t0 = time.time() + # Detailer generation is multi-pass (generate + detect + inpaint per region), use longer timeout + try: + r = requests.post(f'{self.base_url}/sdapi/v1/txt2img', json=payload, timeout=600, verify=False) + if r.status_code != 200: + data = {'error': r.status_code, 'reason': r.reason} + else: + data = r.json() + except requests.exceptions.ConnectionError as e: + self.record(False, 'txt2img_detailer', f"connection error (is a model loaded?): {e}") + return + except requests.exceptions.ReadTimeout: + self.record(False, 'txt2img_detailer', 'timeout after 600s') + return + t1 = time.time() + + if 'error' in data: + self.record(False, 'txt2img_detailer', f"error: {data} (ensure a model is loaded)") + return + + # Should have images + has_images = 'images' in data and len(data['images']) > 0 + self.record(has_images, 'txt2img_detailer_has_images', f"time={t1 - t0:.1f}s") + + if has_images: + # Decode and verify image + from PIL import Image + img_data = data['images'][0].split(',', 1)[0] + img = Image.open(io.BytesIO(base64.b64decode(img_data))) + self.record(True, 'txt2img_detailer_image_valid', f"size={img.size}") + + # Check info field for detailer metadata + if 'info' in data: + info = data['info'] if isinstance(data['info'], str) else json.dumps(data['info']) + has_detailer_info = 'detailer' in info.lower() or 'Detailer' in info + self.record(has_detailer_info, 'txt2img_detailer_metadata', + 'detailer info found in metadata' if has_detailer_info else 'no detailer metadata (detection may have found nothing)') + + def test_txt2img_without_detailer(self): + """POST /sdapi/v1/txt2img baseline without detailer (sanity check).""" + if self._critical_error: + self.skip('txt2img_baseline', self._critical_error) + return + + payload = { + 'prompt': 'a simple landscape', + 'steps': 5, + 'width': 512, + 'height': 512, + 'seed': 42, + 'save_images': False, + 'send_images': True, + } + + data = self._post('/sdapi/v1/txt2img', payload) + if 'error' in data: + self.record(False, 'txt2img_baseline', f"error: {data}") + return + + has_images = 'images' in data and len(data['images']) > 0 + self.record(has_images, 'txt2img_baseline', 'generation works without detailer') + + # ========================================================================= + # Tests: Per-Request Detailer Param Validation + # ========================================================================= + + def _txt2img(self, extra_params=None): + """Helper: generate a portrait with optional param overrides.""" + payload = { + 'prompt': 'a photo of a person, face, portrait', + 'steps': 10, + 'width': 512, + 'height': 512, + 'seed': 42, + 'save_images': False, + 'send_images': True, + } + if extra_params: + payload.update(extra_params) + try: + r = requests.post(f'{self.base_url}/sdapi/v1/txt2img', json=payload, timeout=600, verify=False) + if r.status_code != 200: + return {'error': r.status_code, 'reason': r.reason} + return r.json() + except requests.exceptions.ConnectionError as e: + return {'error': 'connection_refused', 'reason': str(e)} + except requests.exceptions.ReadTimeout: + return {'error': 'timeout', 'reason': 'timeout after 600s'} + + def _decode_image(self, data): + """Decode first image from generation response into numpy array.""" + import numpy as np + from PIL import Image + if 'images' not in data or len(data['images']) == 0: + return None + img_data = data['images'][0].split(',', 1)[0] + img = Image.open(io.BytesIO(base64.b64decode(img_data))).convert('RGB') + return np.array(img, dtype=np.float32) + + def _pixel_diff(self, arr_a, arr_b): + """Mean absolute pixel difference between two images.""" + import numpy as np + if arr_a is None or arr_b is None or arr_a.shape != arr_b.shape: + return -1.0 + return float(np.abs(arr_a - arr_b).mean()) + + def _get_info(self, data): + """Extract info string from generation response.""" + if 'info' not in data: + return '' + info = data['info'] + return info if isinstance(info, str) else json.dumps(info) + + def run_detailer_param_tests(self, available_models=None): + """Verify per-request detailer params change the output.""" + self._category = 'detailer_params' + print("\n--- Per-Request Detailer Param Validation ---") + + if self._critical_error: + self.skip('detailer_params_all', self._critical_error) + return + + # Generate baseline WITHOUT detailer (same seed/prompt as detailer tests) + print(" Generating baseline (no detailer)...") + baseline_data = self._txt2img() + if 'error' in baseline_data: + self.record(False, 'detailer_baseline', f"error: {baseline_data}") + return + baseline = self._decode_image(baseline_data) + if baseline is None: + self.record(False, 'detailer_baseline', 'no image') + return + self.record(True, 'detailer_baseline') + + # Generate WITH detailer enabled (default params) + print(" Generating with detailer (defaults)...") + detailer_default_data = self._txt2img({ + 'detailer_enabled': True, + 'detailer_strength': 0.3, + 'detailer_steps': 5, + 'detailer_conf': 0.3, + }) + if 'error' in detailer_default_data: + self.record(False, 'detailer_default', f"error: {detailer_default_data}") + return + detailer_default = self._decode_image(detailer_default_data) + + # Detailer ON vs OFF should produce different images (if a face was detected) + diff_on_off = self._pixel_diff(baseline, detailer_default) + self.record(diff_on_off > 0.5, 'detailer_on_vs_off', + f"mean_diff={diff_on_off:.2f}" if diff_on_off > 0.5 + else f"identical (diff={diff_on_off:.4f}) — no face detected?") + + # -- Strength variation -- + print(" Testing strength variation...") + strong_data = self._txt2img({ + 'detailer_enabled': True, + 'detailer_strength': 0.7, + 'detailer_steps': 5, + 'detailer_conf': 0.3, + }) + if 'error' not in strong_data: + strong = self._decode_image(strong_data) + diff_strong = self._pixel_diff(detailer_default, strong) + self.record(diff_strong > 0.5, 'detailer_strength_effect', + f"strength 0.3 vs 0.7: diff={diff_strong:.2f}") + + # -- Steps variation -- + print(" Testing steps variation...") + more_steps_data = self._txt2img({ + 'detailer_enabled': True, + 'detailer_strength': 0.3, + 'detailer_steps': 20, + 'detailer_conf': 0.3, + }) + if 'error' not in more_steps_data: + more_steps = self._decode_image(more_steps_data) + diff_steps = self._pixel_diff(detailer_default, more_steps) + self.record(diff_steps > 0.5, 'detailer_steps_effect', + f"steps 5 vs 20: diff={diff_steps:.2f}") + + # -- Resolution variation -- + print(" Testing resolution variation...") + hires_data = self._txt2img({ + 'detailer_enabled': True, + 'detailer_strength': 0.3, + 'detailer_steps': 5, + 'detailer_conf': 0.3, + 'detailer_resolution': 512, + }) + if 'error' not in hires_data: + hires = self._decode_image(hires_data) + diff_res = self._pixel_diff(detailer_default, hires) + self.record(diff_res > 0.5, 'detailer_resolution_effect', + f"resolution 1024 vs 512: diff={diff_res:.2f}") + + # -- Segmentation mode -- + # Segmentation requires a -seg model (e.g. anzhc-face-1024-seg-8n). + # Detection-only models (face-yolo8n) don't produce masks, so the flag has no effect. + seg_models = [m.get('name', '') for m in (available_models or []) + if 'seg' in m.get('name', '').lower() and 'face' in m.get('name', '').lower()] + if seg_models: + seg_model = seg_models[0] + print(f" Testing segmentation mode (model={seg_model})...") + # bbox baseline with the seg model + seg_bbox_data = self._txt2img({ + 'detailer_enabled': True, + 'detailer_strength': 0.3, + 'detailer_steps': 5, + 'detailer_conf': 0.3, + 'detailer_segmentation': False, + 'detailer_models': [seg_model], + }) + seg_data = self._txt2img({ + 'detailer_enabled': True, + 'detailer_strength': 0.3, + 'detailer_steps': 5, + 'detailer_conf': 0.3, + 'detailer_segmentation': True, + 'detailer_models': [seg_model], + }) + if 'error' not in seg_data and 'error' not in seg_bbox_data: + seg_bbox = self._decode_image(seg_bbox_data) + seg_mask = self._decode_image(seg_data) + diff_seg = self._pixel_diff(seg_bbox, seg_mask) + self.record(diff_seg > 0.5, 'detailer_segmentation_effect', + f"bbox vs seg mask ({seg_model}): diff={diff_seg:.2f}") + else: + err = seg_data if 'error' in seg_data else seg_bbox_data + self.record(False, 'detailer_segmentation_effect', f"error: {err}") + else: + print(" Testing segmentation mode...") + seg_data = {'error': 'skipped'} + self.skip('detailer_segmentation_effect', 'no face-seg model available') + + # -- Confidence threshold -- + print(" Testing confidence threshold...") + high_conf_data = self._txt2img({ + 'detailer_enabled': True, + 'detailer_strength': 0.3, + 'detailer_steps': 5, + 'detailer_conf': 0.95, + }) + if 'error' not in high_conf_data: + high_conf = self._decode_image(high_conf_data) + diff_conf = self._pixel_diff(baseline, high_conf) + # High confidence may reject detections, making output closer to baseline + self.record(True, 'detailer_conf_effect', + f"conf=0.95 vs baseline: diff={diff_conf:.2f} " + f"(low diff = detections filtered out, high diff = still detected)") + + # -- Custom detailer prompt -- + print(" Testing detailer prompt override...") + prompt_data = self._txt2img({ + 'detailer_enabled': True, + 'detailer_strength': 0.5, + 'detailer_steps': 5, + 'detailer_conf': 0.3, + 'detailer_prompt': 'a detailed close-up face with freckles', + }) + if 'error' not in prompt_data: + prompt_result = self._decode_image(prompt_data) + diff_prompt = self._pixel_diff(detailer_default, prompt_result) + self.record(diff_prompt > 0.5, 'detailer_prompt_effect', + f"custom prompt vs default: diff={diff_prompt:.2f}") + + # -- Metadata verification across params -- + for test_data, label in [ + (detailer_default_data, 'detailer_default'), + (strong_data if 'error' not in strong_data else None, 'detailer_strong'), + (more_steps_data if 'error' not in more_steps_data else None, 'detailer_more_steps'), + (seg_data if 'error' not in seg_data else None, 'detailer_segmentation'), + ]: + if test_data is None: + continue + info = self._get_info(test_data) + has_meta = 'detailer' in info.lower() or 'Detailer' in info + self.record(has_meta, f'{label}_metadata', + 'detailer info in metadata' if has_meta else 'no detailer metadata') + + # -- Param isolation: generate without detailer after all detailer runs -- + print(" Testing param isolation...") + after_data = self._txt2img() + if 'error' not in after_data: + after = self._decode_image(after_data) + leak_diff = self._pixel_diff(baseline, after) + self.record(leak_diff < 0.5, 'detailer_param_isolation', + f"post-detailer baseline diff={leak_diff:.4f}" if leak_diff < 0.5 + else f"LEAK: baseline changed (diff={leak_diff:.2f})") + + # ========================================================================= + # Runner + # ========================================================================= + + def run_all(self): + print("=" * 60) + print("YOLO Detailer API Test Suite") + print(f"Server: {self.base_url}") + print("=" * 60) + + # Enumerate + models = self.test_detailers_list() + + # Detect across all loaded test images + self.test_detect_all_images(models) + # Test with first available model if any + if models and len(models) > 0: + model_name = models[0].get('name', models[0].get('filename', '')) + if model_name: + self.test_detect_with_model(model_name) + + # Generate + self.test_txt2img_without_detailer() + self.test_txt2img_with_detailer() + + # Per-request detailer param validation + self.run_detailer_param_tests(models) + + # Summary + print("\n" + "=" * 60) + print("Results") + print("=" * 60) + total_passed = 0 + total_failed = 0 + total_skipped = 0 + for cat, data in self.results.items(): + total_passed += data['passed'] + total_failed += data['failed'] + total_skipped += data['skipped'] + status = 'PASS' if data['failed'] == 0 else 'FAIL' + print(f" {cat}: {data['passed']} passed, {data['failed']} failed, {data['skipped']} skipped [{status}]") + print(f" Total: {total_passed} passed, {total_failed} failed, {total_skipped} skipped") + print("=" * 60) + return total_failed == 0 + + +if __name__ == "__main__": + parser = argparse.ArgumentParser(description='YOLO Detailer API Tests') + parser.add_argument('--url', default=os.environ.get('SDAPI_URL', 'http://127.0.0.1:7860'), help='server URL') + parser.add_argument('--image', default=None, help='test image path') + args = parser.parse_args() + test = DetailerAPITest(args.url, args.image) + success = test.run_all() + sys.exit(0 if success else 1) diff --git a/test/test-generation-api.py b/test/test-generation-api.py new file mode 100644 index 000000000..1359d59e5 --- /dev/null +++ b/test/test-generation-api.py @@ -0,0 +1,615 @@ +#!/usr/bin/env python +""" +API tests for generation with scheduler params, color grading, and latent corrections. + +Tests: +- GET /sdapi/v1/samplers — sampler enumeration and config +- POST /sdapi/v1/txt2img — generation with various samplers +- POST /sdapi/v1/txt2img — generation with color grading params +- POST /sdapi/v1/txt2img — generation with latent correction params + +Requires a running SD.Next instance with a model loaded. + +Usage: + python test/test-generation-api.py [--url URL] [--steps STEPS] +""" + +import io +import os +import sys +import json +import time +import base64 +import argparse +import requests +import urllib3 + +urllib3.disable_warnings(urllib3.exceptions.InsecureRequestWarning) + + +class GenerationAPITest: + """Test harness for generation API with scheduler and grading params.""" + + # Samplers to test — a representative subset covering different scheduler families + TEST_SAMPLERS = [ + 'Euler a', + 'Euler', + 'DPM++ 2M', + 'UniPC', + 'DDIM', + 'DPM++ 2M SDE', + ] + + def __init__(self, base_url, steps=10, timeout=300): + self.base_url = base_url.rstrip('/') + self.steps = steps + self.timeout = timeout + self.results = { + 'samplers': {'passed': 0, 'failed': 0, 'skipped': 0, 'tests': []}, + 'generation': {'passed': 0, 'failed': 0, 'skipped': 0, 'tests': []}, + 'grading': {'passed': 0, 'failed': 0, 'skipped': 0, 'tests': []}, + 'correction': {'passed': 0, 'failed': 0, 'skipped': 0, 'tests': []}, + 'param_validation': {'passed': 0, 'failed': 0, 'skipped': 0, 'tests': []}, + } + self._category = 'samplers' + self._critical_error = None + + def _get(self, endpoint): + try: + r = requests.get(f'{self.base_url}{endpoint}', timeout=self.timeout, verify=False) + if r.status_code != 200: + return {'error': r.status_code, 'reason': r.reason} + return r.json() + except requests.exceptions.ConnectionError: + return {'error': 'connection_refused', 'reason': 'Server not running'} + except Exception as e: + return {'error': 'exception', 'reason': str(e)} + + def _post(self, endpoint, data): + try: + r = requests.post(f'{self.base_url}{endpoint}', json=data, timeout=self.timeout, verify=False) + if r.status_code != 200: + return {'error': r.status_code, 'reason': r.reason} + return r.json() + except requests.exceptions.ConnectionError: + return {'error': 'connection_refused', 'reason': 'Server not running'} + except Exception as e: + return {'error': 'exception', 'reason': str(e)} + + def record(self, passed, name, detail=''): + status = 'PASS' if passed else 'FAIL' + self.results[self._category]['passed' if passed else 'failed'] += 1 + self.results[self._category]['tests'].append((status, name)) + msg = f' {status}: {name}' + if detail: + msg += f' ({detail})' + print(msg) + + def skip(self, name, reason): + self.results[self._category]['skipped'] += 1 + self.results[self._category]['tests'].append(('SKIP', name)) + print(f' SKIP: {name} ({reason})') + + def _txt2img(self, extra_params=None, prompt='a cat'): + """Helper: run txt2img with base params + overrides. Returns (data, time).""" + payload = { + 'prompt': prompt, + 'steps': self.steps, + 'width': 512, + 'height': 512, + 'seed': 42, + 'save_images': False, + 'send_images': True, + } + if extra_params: + payload.update(extra_params) + t0 = time.time() + data = self._post('/sdapi/v1/txt2img', payload) + return data, time.time() - t0 + + def _check_generation(self, data, test_name, elapsed): + """Validate a generation response has images.""" + if 'error' in data: + self.record(False, test_name, f"error: {data}") + return False + has_images = 'images' in data and len(data['images']) > 0 + self.record(has_images, test_name, f"time={elapsed:.1f}s") + return has_images + + def _get_info(self, data): + """Extract info string from generation response.""" + if 'info' not in data: + return '' + info = data['info'] + return info if isinstance(info, str) else json.dumps(info) + + def _decode_image(self, data): + """Decode first image from generation response into numpy array.""" + import numpy as np + from PIL import Image + if 'images' not in data or len(data['images']) == 0: + return None + img_data = data['images'][0].split(',', 1)[0] + img = Image.open(io.BytesIO(base64.b64decode(img_data))).convert('RGB') + return np.array(img, dtype=np.float32) + + def _pixel_diff(self, arr_a, arr_b): + """Mean absolute pixel difference between two images (0-255 scale).""" + import numpy as np + if arr_a is None or arr_b is None: + return -1.0 + if arr_a.shape != arr_b.shape: + return -1.0 + return float(np.abs(arr_a - arr_b).mean()) + + def _channel_means(self, arr): + """Return per-channel means [R, G, B].""" + if arr is None: + return [0, 0, 0] + return [float(arr[:, :, c].mean()) for c in range(3)] + + # ========================================================================= + # Tests: Sampler Enumeration + # ========================================================================= + + def test_samplers_list(self): + """GET /sdapi/v1/samplers returns available samplers with config.""" + self._category = 'samplers' + print("\n--- Sampler Enumeration ---") + + data = self._get('/sdapi/v1/samplers') + if 'error' in data: + self.record(False, 'samplers_list', f"error: {data}") + self._critical_error = f"Server error: {data}" + return [] + + if not isinstance(data, list): + self.record(False, 'samplers_list', f"expected list, got {type(data).__name__}") + return [] + + self.record(True, 'samplers_list', f"{len(data)} samplers available") + + # Check that each sampler has a name + sampler_names = [] + for s in data: + name = s.get('name', '') + if name: + sampler_names.append(name) + + self.record(len(sampler_names) == len(data), 'samplers_have_names', + f"{len(sampler_names)}/{len(data)} have names") + + # Check for our test samplers + for test_sampler in self.TEST_SAMPLERS: + found = test_sampler in sampler_names + if not found: + self.skip(f'sampler_available_{test_sampler}', 'not in server sampler list') + else: + self.record(True, f'sampler_available_{test_sampler}') + + return sampler_names + + # ========================================================================= + # Tests: Generation with Different Samplers + # ========================================================================= + + def test_samplers_generate(self, available_samplers): + """Generate with each test sampler and verify success.""" + self._category = 'generation' + print("\n--- Generation with Different Samplers ---") + + if self._critical_error: + for s in self.TEST_SAMPLERS: + self.skip(f'generate_{s}', self._critical_error) + return + + for sampler in self.TEST_SAMPLERS: + if sampler not in available_samplers: + self.skip(f'generate_{sampler}', 'sampler not available') + continue + data, elapsed = self._txt2img({'sampler_name': sampler}) + self._check_generation(data, f'generate_{sampler}', elapsed) + + # ========================================================================= + # Tests: Color Grading Params + # ========================================================================= + + def test_grading_brightness_contrast(self): + """Generate with grading brightness and contrast.""" + data, elapsed = self._txt2img({ + 'grading_brightness': 0.2, + 'grading_contrast': 0.3, + }) + self._check_generation(data, 'grading_brightness_contrast', elapsed) + + def test_grading_saturation_hue(self): + """Generate with grading saturation and hue shift.""" + data, elapsed = self._txt2img({ + 'grading_saturation': 0.5, + 'grading_hue': 0.1, + }) + self._check_generation(data, 'grading_saturation_hue', elapsed) + + def test_grading_gamma_sharpness(self): + """Generate with gamma correction and sharpness.""" + data, elapsed = self._txt2img({ + 'grading_gamma': 0.8, + 'grading_sharpness': 0.5, + }) + self._check_generation(data, 'grading_gamma_sharpness', elapsed) + + def test_grading_color_temp(self): + """Generate with warm color temperature.""" + data, elapsed = self._txt2img({ + 'grading_color_temp': 3500, + }) + self._check_generation(data, 'grading_color_temp', elapsed) + + def test_grading_tone(self): + """Generate with shadows/midtones/highlights adjustments.""" + data, elapsed = self._txt2img({ + 'grading_shadows': 0.3, + 'grading_midtones': -0.1, + 'grading_highlights': 0.2, + }) + self._check_generation(data, 'grading_tone', elapsed) + + def test_grading_effects(self): + """Generate with vignette and grain.""" + data, elapsed = self._txt2img({ + 'grading_vignette': 0.5, + 'grading_grain': 0.3, + }) + self._check_generation(data, 'grading_effects', elapsed) + + def test_grading_split_toning(self): + """Generate with split toning colors.""" + data, elapsed = self._txt2img({ + 'grading_shadows_tint': '#003366', + 'grading_highlights_tint': '#ffcc00', + 'grading_split_tone_balance': 0.6, + }) + self._check_generation(data, 'grading_split_toning', elapsed) + + def test_grading_combined(self): + """Generate with multiple grading params at once.""" + data, elapsed = self._txt2img({ + 'grading_brightness': 0.1, + 'grading_contrast': 0.2, + 'grading_saturation': 0.3, + 'grading_gamma': 0.9, + 'grading_color_temp': 5000, + 'grading_vignette': 0.3, + }) + self._check_generation(data, 'grading_combined', elapsed) + + def run_grading_tests(self): + """Run all grading tests.""" + self._category = 'grading' + print("\n--- Color Grading Tests ---") + + if self._critical_error: + self.skip('grading_all', self._critical_error) + return + + self.test_grading_brightness_contrast() + self.test_grading_saturation_hue() + self.test_grading_gamma_sharpness() + self.test_grading_color_temp() + self.test_grading_tone() + self.test_grading_effects() + self.test_grading_split_toning() + self.test_grading_combined() + + # ========================================================================= + # Tests: Latent Correction Params + # ========================================================================= + + def test_correction_brightness(self): + """Generate with latent brightness correction.""" + data, elapsed = self._txt2img({'hdr_brightness': 1.5}) + ok = self._check_generation(data, 'correction_brightness', elapsed) + if ok: + info = self._get_info(data) + has_param = 'Latent brightness' in info + self.record(has_param, 'correction_brightness_metadata', + 'found in info' if has_param else 'not found in info') + + def test_correction_color(self): + """Generate with latent color centering.""" + data, elapsed = self._txt2img({'hdr_color': 0.5, 'hdr_mode': 1}) + ok = self._check_generation(data, 'correction_color', elapsed) + if ok: + info = self._get_info(data) + has_param = 'Latent color' in info + self.record(has_param, 'correction_color_metadata', + 'found in info' if has_param else 'not found in info') + + def test_correction_clamp(self): + """Generate with latent clamping.""" + data, elapsed = self._txt2img({ + 'hdr_clamp': True, + 'hdr_threshold': 0.8, + 'hdr_boundary': 4.0, + }) + ok = self._check_generation(data, 'correction_clamp', elapsed) + if ok: + info = self._get_info(data) + has_param = 'Latent clamp' in info + self.record(has_param, 'correction_clamp_metadata', + 'found in info' if has_param else 'not found in info') + + def test_correction_sharpen(self): + """Generate with latent sharpening.""" + data, elapsed = self._txt2img({'hdr_sharpen': 1.0}) + ok = self._check_generation(data, 'correction_sharpen', elapsed) + if ok: + info = self._get_info(data) + has_param = 'Latent sharpen' in info + self.record(has_param, 'correction_sharpen_metadata', + 'found in info' if has_param else 'not found in info') + + def test_correction_maximize(self): + """Generate with latent maximize/normalize.""" + data, elapsed = self._txt2img({ + 'hdr_maximize': True, + 'hdr_max_center': 0.6, + 'hdr_max_boundary': 2.0, + }) + ok = self._check_generation(data, 'correction_maximize', elapsed) + if ok: + info = self._get_info(data) + has_param = 'Latent max' in info + self.record(has_param, 'correction_maximize_metadata', + 'found in info' if has_param else 'not found in info') + + def test_correction_combined(self): + """Generate with multiple correction params.""" + data, elapsed = self._txt2img({ + 'hdr_brightness': 1.0, + 'hdr_color': 0.3, + 'hdr_sharpen': 0.5, + 'hdr_clamp': True, + }) + ok = self._check_generation(data, 'correction_combined', elapsed) + if ok: + info = self._get_info(data) + # At least some correction params should appear + found = [k for k in ['Latent brightness', 'Latent color', 'Latent sharpen', 'Latent clamp'] if k in info] + self.record(len(found) > 0, 'correction_combined_metadata', f"found: {found}") + + def run_correction_tests(self): + """Run all latent correction tests.""" + self._category = 'correction' + print("\n--- Latent Correction Tests ---") + + if self._critical_error: + self.skip('correction_all', self._critical_error) + return + + self.test_correction_brightness() + self.test_correction_color() + self.test_correction_clamp() + self.test_correction_sharpen() + self.test_correction_maximize() + self.test_correction_combined() + + # ========================================================================= + # Tests: Per-Request Param Validation (baseline comparison) + # ========================================================================= + + def _generate_baseline(self): + """Generate a baseline image with no grading/correction. Cache and reuse.""" + if hasattr(self, '_baseline_arr') and self._baseline_arr is not None: + return self._baseline_arr, self._baseline_data + data, elapsed = self._txt2img() + if 'error' in data or 'images' not in data: + return None, data + self._baseline_arr = self._decode_image(data) + self._baseline_data = data + print(f' Baseline generated: time={elapsed:.1f}s mean={self._channel_means(self._baseline_arr)}') + return self._baseline_arr, data + + def _compare_param(self, name, params, check_fn=None): + """Generate with params and compare to baseline. Optionally run check_fn(baseline, result).""" + baseline, _ = self._generate_baseline() + if baseline is None: + self.skip(f'param_{name}', 'baseline generation failed') + return + + data, elapsed = self._txt2img(params) + if 'error' in data: + self.record(False, f'param_{name}', f"generation error: {data}") + return + + result = self._decode_image(data) + if result is None: + self.record(False, f'param_{name}', 'no image in response') + return + + diff = self._pixel_diff(baseline, result) + differs = diff > 0.5 # more than 0.5/255 mean difference + self.record(differs, f'param_{name}_differs', + f"mean_diff={diff:.2f}" if differs else f"images identical (diff={diff:.4f})") + + if check_fn and differs: + try: + ok, detail = check_fn(baseline, result, data) + self.record(ok, f'param_{name}_direction', detail) + except Exception as e: + self.record(False, f'param_{name}_direction', f"check error: {e}") + + def run_param_validation_tests(self): + """Verify per-request grading/correction params actually change the output.""" + self._category = 'param_validation' + print("\n--- Per-Request Param Validation ---") + + if self._critical_error: + self.skip('param_validation_all', self._critical_error) + return + + import numpy as np + + # -- Grading params -- + + # Brightness: positive should increase mean pixel value + def check_brightness(base, result, _data): + base_mean = float(base.mean()) + result_mean = float(result.mean()) + return result_mean > base_mean, f"baseline={base_mean:.1f} graded={result_mean:.1f}" + self._compare_param('grading_brightness', {'grading_brightness': 0.3}, check_brightness) + + # Contrast: should increase standard deviation + def check_contrast(base, result, _data): + return float(result.std()) > float(base.std()), \ + f"baseline_std={float(base.std()):.1f} graded_std={float(result.std()):.1f}" + self._compare_param('grading_contrast', {'grading_contrast': 0.5}, check_contrast) + + # Saturation: desaturation should reduce color channel spread + def check_desaturation(base, result, _data): + base_spread = max(self._channel_means(base)) - min(self._channel_means(base)) + result_spread = max(self._channel_means(result)) - min(self._channel_means(result)) + return result_spread < base_spread, \ + f"baseline_spread={base_spread:.1f} graded_spread={result_spread:.1f}" + self._compare_param('grading_saturation_neg', {'grading_saturation': -0.5}, check_desaturation) + + # Hue shift: just verify it changes + self._compare_param('grading_hue', {'grading_hue': 0.2}) + + # Gamma < 1: should brighten (raise values that are < 1) + def check_gamma(base, result, _data): + return float(result.mean()) > float(base.mean()), \ + f"baseline={float(base.mean()):.1f} gamma={float(result.mean()):.1f}" + self._compare_param('grading_gamma', {'grading_gamma': 0.7}, check_gamma) + + # Sharpness: just verify it changes + self._compare_param('grading_sharpness', {'grading_sharpness': 0.8}) + + # Color temperature warm: red channel mean should increase relative to blue + def check_warm(base, result, _data): + base_r, _, base_b = self._channel_means(base) + res_r, _, res_b = self._channel_means(result) + base_rb = base_r - base_b + res_rb = res_r - res_b + return res_rb > base_rb, f"baseline R-B={base_rb:.1f} warm R-B={res_rb:.1f}" + self._compare_param('grading_color_temp_warm', {'grading_color_temp': 3000}, check_warm) + + # Color temperature cool: blue should increase relative to red + def check_cool(base, result, _data): + base_r, _, base_b = self._channel_means(base) + res_r, _, res_b = self._channel_means(result) + base_rb = base_r - base_b + res_rb = res_r - res_b + return res_rb < base_rb, f"baseline R-B={base_rb:.1f} cool R-B={res_rb:.1f}" + self._compare_param('grading_color_temp_cool', {'grading_color_temp': 10000}, check_cool) + + # Vignette: corners should be darker than baseline corners + def check_vignette(base, result, _data): + h, w = base.shape[:2] + corner_size = h // 8 + base_corners = np.concatenate([ + base[:corner_size, :corner_size].flatten(), + base[:corner_size, -corner_size:].flatten(), + base[-corner_size:, :corner_size].flatten(), + base[-corner_size:, -corner_size:].flatten(), + ]) + result_corners = np.concatenate([ + result[:corner_size, :corner_size].flatten(), + result[:corner_size, -corner_size:].flatten(), + result[-corner_size:, :corner_size].flatten(), + result[-corner_size:, -corner_size:].flatten(), + ]) + return float(result_corners.mean()) < float(base_corners.mean()), \ + f"baseline_corners={float(base_corners.mean()):.1f} vignette_corners={float(result_corners.mean()):.1f}" + self._compare_param('grading_vignette', {'grading_vignette': 0.8}, check_vignette) + + # Grain: just verify it changes (stochastic) + self._compare_param('grading_grain', {'grading_grain': 0.5}) + + # Shadows/midtones/highlights: verify changes + self._compare_param('grading_shadows', {'grading_shadows': 0.5}) + self._compare_param('grading_highlights', {'grading_highlights': -0.3}) + + # CLAHE: should increase local contrast + self._compare_param('grading_clahe', {'grading_clahe_clip': 2.0}) + + # Split toning: verify changes + self._compare_param('grading_split_toning', { + 'grading_shadows_tint': '#003366', + 'grading_highlights_tint': '#ffcc00', + }) + + # -- Correction params -- + + # Latent brightness: should change output and appear in metadata + def check_correction_meta(key): + def _check(_base, _result, data): + info = self._get_info(data) + return key in info, f"'{key}' {'found' if key in info else 'missing'} in info" + return _check + self._compare_param('hdr_brightness', {'hdr_brightness': 2.0}, check_correction_meta('Latent brightness')) + self._compare_param('hdr_color', {'hdr_color': 0.8, 'hdr_mode': 1}, check_correction_meta('Latent color')) + self._compare_param('hdr_sharpen', {'hdr_sharpen': 1.5}, check_correction_meta('Latent sharpen')) + self._compare_param('hdr_clamp', {'hdr_clamp': True, 'hdr_threshold': 0.7}, check_correction_meta('Latent clamp')) + + # Isolation: verify params from one request don't leak to the next + data_after, _ = self._txt2img() + arr_after = self._decode_image(data_after) + baseline, _ = self._generate_baseline() + if baseline is not None and arr_after is not None: + leak_diff = self._pixel_diff(baseline, arr_after) + no_leak = leak_diff < 0.5 + self.record(no_leak, 'param_isolation', + f"post-grading baseline diff={leak_diff:.4f}" if no_leak + else f"LEAK: baseline changed after grading requests (diff={leak_diff:.2f})") + + # ========================================================================= + # Runner + # ========================================================================= + + def run_all(self): + print("=" * 60) + print("Generation API Test Suite") + print(f"Server: {self.base_url}") + print(f"Steps: {self.steps}") + print("=" * 60) + + # Samplers + available = self.test_samplers_list() + self.test_samplers_generate(available) + + # Grading + self.run_grading_tests() + + # Corrections + self.run_correction_tests() + + # Per-request param validation (baseline comparison) + self.run_param_validation_tests() + + # Summary + print("\n" + "=" * 60) + print("Results") + print("=" * 60) + total_passed = 0 + total_failed = 0 + total_skipped = 0 + for cat, data in self.results.items(): + total_passed += data['passed'] + total_failed += data['failed'] + total_skipped += data['skipped'] + status = 'PASS' if data['failed'] == 0 else 'FAIL' + print(f" {cat}: {data['passed']} passed, {data['failed']} failed, {data['skipped']} skipped [{status}]") + print(f" Total: {total_passed} passed, {total_failed} failed, {total_skipped} skipped") + print("=" * 60) + return total_failed == 0 + + +if __name__ == "__main__": + parser = argparse.ArgumentParser(description='Generation API Tests (samplers, grading, correction)') + parser.add_argument('--url', default=os.environ.get('SDAPI_URL', 'http://127.0.0.1:7860'), help='server URL') + parser.add_argument('--steps', type=int, default=10, help='generation steps (lower = faster tests)') + args = parser.parse_args() + test = GenerationAPITest(args.url, args.steps) + success = test.run_all() + sys.exit(0 if success else 1) diff --git a/test/test-grading.py b/test/test-grading.py new file mode 100644 index 000000000..2018519c6 --- /dev/null +++ b/test/test-grading.py @@ -0,0 +1,633 @@ +#!/usr/bin/env python +""" +Offline unit tests for color grading and latent corrections. + +Tests two systems: +- Pixel-space color grading (modules/processing_grading.py) +- Latent-space corrections (modules/processing_correction.py) + +No running server required. Tests core logic with synthetic inputs. + +Usage: + python test/test-grading.py +""" + +import os +import sys +import time +import types +import torch +import numpy as np +from types import SimpleNamespace + +script_dir = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) +sys.path.insert(0, script_dir) +os.chdir(script_dir) + +os.environ['SD_INSTALL_QUIET'] = '1' + +# Initialize cmd_args before any module imports (required by shared.py) +import modules.cmd_args +import installer +installer.add_args(modules.cmd_args.parser) +modules.cmd_args.parsed, _ = modules.cmd_args.parser.parse_known_args([]) + +# Mock sd_vae_taesd to break circular import: +# processing_correction -> sd_vae_taesd -> shared -> shared_items -> sd_vae_taesd (circle) +_mock_taesd = types.ModuleType('modules.vae.sd_vae_taesd') +_mock_taesd.TAESD_MODELS = {'taesd': None} +_mock_taesd.CQYAN_MODELS = {} +_mock_taesd.encode = lambda x: torch.zeros(1, 4, 8, 8) +sys.modules['modules.vae.sd_vae_taesd'] = _mock_taesd + +from modules.errors import log + +# Results tracking +results = { + 'grading_params': {'passed': 0, 'failed': 0, 'tests': []}, + 'grading_functions': {'passed': 0, 'failed': 0, 'tests': []}, + 'correction_primitives': {'passed': 0, 'failed': 0, 'tests': []}, + 'correction_pipeline': {'passed': 0, 'failed': 0, 'tests': []}, +} +current_category = 'grading_params' + + +def record(passed, name, detail=''): + status = 'PASS' if passed else 'FAIL' + results[current_category]['passed' if passed else 'failed'] += 1 + results[current_category]['tests'].append((status, name)) + msg = f' {status}: {name}' + if detail: + msg += f' ({detail})' + if passed: + log.info(msg) + else: + log.error(msg) + + +def set_category(cat): + global current_category # pylint: disable=global-statement + current_category = cat + + +# ============================================================ +# Color Grading: GradingParams and utility functions +# ============================================================ + +def test_grading_params_defaults(): + """GradingParams has correct defaults.""" + from modules.processing_grading import GradingParams + p = GradingParams() + assert p.brightness == 0.0 + assert p.contrast == 0.0 + assert p.saturation == 0.0 + assert p.hue == 0.0 + assert p.gamma == 1.0 + assert p.sharpness == 0.0 + assert p.color_temp == 6500 + assert p.shadows == 0.0 + assert p.midtones == 0.0 + assert p.highlights == 0.0 + assert p.clahe_clip == 0.0 + assert p.clahe_grid == 8 + assert p.shadows_tint == "#000000" + assert p.highlights_tint == "#ffffff" + assert p.split_tone_balance == 0.5 + assert p.vignette == 0.0 + assert p.grain == 0.0 + assert p.lut_file == "" + assert p.lut_strength == 1.0 + return True + + +def test_grading_is_active(): + """is_active() returns False for defaults, True when any param differs.""" + from modules.processing_grading import GradingParams, is_active + assert not is_active(GradingParams()), "defaults should be inactive" + assert is_active(GradingParams(brightness=0.1)), "non-default brightness should be active" + assert is_active(GradingParams(gamma=0.9)), "non-default gamma should be active" + assert is_active(GradingParams(shadows_tint="#ff0000")), "non-default tint should be active" + assert is_active(GradingParams(vignette=0.5)), "non-default vignette should be active" + assert not is_active(GradingParams(brightness=0.0, gamma=1.0, color_temp=6500)), "all-default should be inactive" + return True + + +def test_grading_float_coercion(): + """__post_init__ coerces int inputs to float (Gradio sends int for float sliders).""" + from modules.processing_grading import GradingParams + p = GradingParams(brightness=1, contrast=2, gamma=1, color_temp=6500) + assert isinstance(p.brightness, float), f"expected float, got {type(p.brightness)}" + assert isinstance(p.contrast, float), f"expected float, got {type(p.contrast)}" + assert isinstance(p.gamma, float), f"expected float, got {type(p.gamma)}" + assert isinstance(p.color_temp, float), f"expected float, got {type(p.color_temp)}" + return True + + +def test_hex_to_rgb(): + """_hex_to_rgb converts hex color strings correctly.""" + from modules.processing_grading import _hex_to_rgb + assert _hex_to_rgb("#000000") == (0.0, 0.0, 0.0), "black" + assert _hex_to_rgb("#ffffff") == (1.0, 1.0, 1.0), "white" + r, g, b = _hex_to_rgb("#ff0000") + assert abs(r - 1.0) < 1e-6 and abs(g) < 1e-6 and abs(b) < 1e-6, "red" + r, g, b = _hex_to_rgb("#00ff00") + assert abs(r) < 1e-6 and abs(g - 1.0) < 1e-6 and abs(b) < 1e-6, "green" + r, g, b = _hex_to_rgb("#0000ff") + assert abs(r) < 1e-6 and abs(g) < 1e-6 and abs(b - 1.0) < 1e-6, "blue" + # without hash + r, g, b = _hex_to_rgb("ff8040") + assert r > g > b, "orange-ish ordering" + # invalid length returns black + assert _hex_to_rgb("#fff") == (0.0, 0.0, 0.0), "short hex returns black" + return True + + +def test_kelvin_to_rgb(): + """_kelvin_to_rgb_scale returns sensible values at known temperatures.""" + from modules.processing_grading import _kelvin_to_rgb_scale + # 6500K (reference) should be approximately (1, 1, 1) - tolerance is wide because + # the Planckian formula approximation normalizes to a hardcoded ref point + r, g, b = _kelvin_to_rgb_scale(6500) + assert abs(r - 1.0) < 0.15 and abs(g - 1.0) < 0.15 and abs(b - 1.0) < 0.15, f"6500K: ({r:.3f}, {g:.3f}, {b:.3f})" + # warm (3000K) should have r > b + r, g, b = _kelvin_to_rgb_scale(3000) + assert r > b, f"3000K should be warm: r={r:.3f} b={b:.3f}" + # cool (10000K) should have b > r + r, g, b = _kelvin_to_rgb_scale(10000) + assert b > r, f"10000K should be cool: r={r:.3f} b={b:.3f}" + # all values should be non-negative; very low temps have zero blue (physically correct) + for temp in [1000, 2000, 4000, 8000, 15000, 40000]: + r, g, b = _kelvin_to_rgb_scale(temp) + assert r >= 0 and g >= 0 and b >= 0, f"{temp}K has negative channel: ({r:.3f}, {g:.3f}, {b:.3f})" + # moderate temps should have all positive channels + for temp in [3000, 5000, 6500, 10000]: + r, g, b = _kelvin_to_rgb_scale(temp) + assert r > 0 and g > 0 and b > 0, f"{temp}K has non-positive channel: ({r:.3f}, {g:.3f}, {b:.3f})" + return True + + +# ============================================================ +# Color Grading: torch-based functions (need kornia for some) +# ============================================================ + +def _make_test_image_tensor(h=64, w=64): + """Create a synthetic RGB test image tensor [1, 3, H, W] in [0, 1].""" + torch.manual_seed(42) + return torch.rand(1, 3, h, w, dtype=torch.float32) + + +def _make_test_pil_image(h=64, w=64): + """Create a synthetic RGB PIL image.""" + from PIL import Image + arr = np.random.RandomState(42).randint(0, 255, (h, w, 3), dtype=np.uint8) + return Image.fromarray(arr, 'RGB') + + +def test_apply_vignette(): + """_apply_vignette darkens edges more than center.""" + from modules.processing_grading import _apply_vignette + img = torch.ones(1, 3, 64, 64, dtype=torch.float32) + result = _apply_vignette(img, strength=1.0) + center_val = result[0, 0, 32, 32].item() + corner_val = result[0, 0, 0, 0].item() + assert center_val > corner_val, f"center ({center_val:.3f}) should be brighter than corner ({corner_val:.3f})" + assert result.shape == img.shape, "shape preserved" + assert not torch.isnan(result).any(), "no NaN" + # zero strength should be identity + result_zero = _apply_vignette(img, strength=0.0) + assert torch.allclose(result_zero, img), "zero strength is identity" + return True + + +def test_apply_grain(): + """_apply_grain adds noise (output differs from input, stays in valid range).""" + from modules.processing_grading import _apply_grain + img = torch.ones(1, 3, 64, 64, dtype=torch.float32) * 0.5 + result = _apply_grain(img, strength=0.5) + assert not torch.equal(result, img), "grain should modify the image" + assert result.shape == img.shape, "shape preserved" + assert result.min() >= 0.0 and result.max() <= 1.0, "output clamped to [0, 1]" + assert not torch.isnan(result).any(), "no NaN" + return True + + +def test_apply_color_temp(): + """_apply_color_temp shifts R/B channels for warm/cool temperatures.""" + from modules.processing_grading import _apply_color_temp + img = torch.ones(1, 3, 64, 64, dtype=torch.float32) * 0.5 + # warm + warm = _apply_color_temp(img, 3000) + assert warm[0, 0].mean() > warm[0, 2].mean(), "warm: red > blue" + # cool + cool = _apply_color_temp(img, 10000) + assert cool[0, 2].mean() > cool[0, 0].mean(), "cool: blue > red" + # neutral + neutral = _apply_color_temp(img, 6500) + assert torch.allclose(neutral, img, atol=0.05), "6500K is near-neutral" + assert warm.shape == img.shape, "shape preserved" + return True + + +def test_apply_shadows_midtones_highlights(): + """_apply_shadows_midtones_highlights modifies tone without NaN/shape issues.""" + try: + from modules.processing_grading import _apply_shadows_midtones_highlights + except ImportError: + return None # kornia not available + img = _make_test_image_tensor() + # shadows boost + result = _apply_shadows_midtones_highlights(img, shadows=0.5, midtones=0.0, highlights=0.0) + assert result.shape == img.shape, "shape preserved" + assert not torch.isnan(result).any(), "no NaN" + assert result.min() >= 0.0 and result.max() <= 1.0, "output in [0, 1]" + # all zero should be near-identity (kornia conversions may introduce tiny diffs) + result_zero = _apply_shadows_midtones_highlights(img, shadows=0.0, midtones=0.0, highlights=0.0) + assert torch.allclose(result_zero, img, atol=1e-3), "zero params is near-identity" + return True + + +def test_grade_image_pipeline(): + """Full grade_image pipeline runs without errors for various param combos.""" + try: + import modules.devices as devices_mod + devices_mod.device = torch.device('cpu') + devices_mod.dtype = torch.float32 + from modules.processing_grading import GradingParams, grade_image, is_active + except ImportError: + return None # kornia not available + img = _make_test_pil_image() + # basic adjustments + params = GradingParams(brightness=0.1, contrast=0.2, saturation=-0.1) + assert is_active(params) + result = grade_image(img, params) + assert result.size == img.size, "output size matches input" + assert result.mode == 'RGB', "output is RGB" + # tone adjustments + params = GradingParams(shadows=0.3, midtones=-0.2, highlights=0.1) + result = grade_image(img, params) + assert result.size == img.size + # effects + params = GradingParams(vignette=0.5, grain=0.3) + result = grade_image(img, params) + assert result.size == img.size + # hue and gamma + params = GradingParams(hue=0.1, gamma=0.8, sharpness=0.5) + result = grade_image(img, params) + assert result.size == img.size + # color temp + params = GradingParams(color_temp=3000) + result = grade_image(img, params) + assert result.size == img.size + # split toning + params = GradingParams(shadows_tint="#003366", highlights_tint="#ffcc00", split_tone_balance=0.7) + result = grade_image(img, params) + assert result.size == img.size + return True + + +def test_grade_image_edge_cases(): + """grade_image handles edge cases: all-black, all-white, tiny images.""" + try: + import modules.devices as devices_mod + devices_mod.device = torch.device('cpu') + devices_mod.dtype = torch.float32 + from modules.processing_grading import GradingParams, grade_image + except ImportError: + return None + from PIL import Image + params = GradingParams(brightness=0.2, contrast=0.3, vignette=0.5, grain=0.2) + # all black + black = Image.fromarray(np.zeros((64, 64, 3), dtype=np.uint8), 'RGB') + result = grade_image(black, params) + assert result.size == black.size, "black image handled" + # all white + white = Image.fromarray(np.full((64, 64, 3), 255, dtype=np.uint8), 'RGB') + result = grade_image(white, params) + assert result.size == white.size, "white image handled" + # tiny image + tiny = Image.fromarray(np.random.randint(0, 255, (4, 4, 3), dtype=np.uint8), 'RGB') + result = grade_image(tiny, params) + assert result.size == tiny.size, "tiny image handled" + return True + + +# ============================================================ +# Latent Corrections: primitive tensor operations +# ============================================================ + +def test_soft_clamp_tensor(): + """soft_clamp_tensor shrinks outliers toward mean, preserves values within bounds.""" + from modules.processing_correction import soft_clamp_tensor + # within bounds: no change + tensor = torch.randn(4, 64, 64) * 0.5 + result = soft_clamp_tensor(tensor, threshold=0.8, boundary=4) + assert torch.allclose(result, tensor, atol=1e-5), "within-bounds tensor unchanged" + # with outliers: should clamp + tensor_outliers = torch.randn(4, 64, 64) + tensor_outliers[0, 0, 0] = 10.0 + tensor_outliers[1, 0, 0] = -10.0 + result = soft_clamp_tensor(tensor_outliers.clone(), threshold=0.8, boundary=4) + assert result[0, 0, 0] < tensor_outliers[0, 0, 0], "positive outlier reduced" + assert result[1, 0, 0] > tensor_outliers[1, 0, 0], "negative outlier raised" + assert result.shape == tensor_outliers.shape, "shape preserved" + assert not torch.isnan(result).any(), "no NaN" + # zero threshold: identity + result_zero = soft_clamp_tensor(tensor_outliers.clone(), threshold=0, boundary=4) + assert torch.allclose(result_zero, tensor_outliers), "zero threshold is identity" + return True + + +def test_center_tensor(): + """center_tensor adjusts mean of tensor channels.""" + from modules.processing_correction import center_tensor + tensor = torch.randn(4, 64, 64) + 2.0 # offset mean + original_mean = tensor.mean().item() + # full shift should reduce mean toward offset + result = center_tensor(tensor.clone(), channel_shift=0.0, full_shift=1.0, offset=0.0) + assert abs(result.mean().item()) < abs(original_mean), "full shift centers toward zero" + # channel shift + result_ch = center_tensor(tensor.clone(), channel_shift=1.0, full_shift=0.0, offset=0.0) + for c in range(4): + assert abs(result_ch[c].mean().item()) < abs(tensor[c].mean().item()), f"channel {c} centered" + # no-op + result_noop = center_tensor(tensor.clone(), channel_shift=0.0, full_shift=0.0, offset=0.0) + assert torch.allclose(result_noop, tensor), "zero params is identity" + # with offset + result_offset = center_tensor(tensor.clone(), channel_shift=0.0, full_shift=1.0, offset=5.0) + assert result_offset.mean().item() > 0, "offset shifts mean positive" + return True + + +def test_sharpen_tensor(): + """sharpen_tensor applies sharpening convolution, preserves shape.""" + from modules.processing_correction import sharpen_tensor + tensor = torch.randn(4, 64, 64) + # zero ratio: identity + result_zero = sharpen_tensor(tensor.clone(), ratio=0) + assert torch.allclose(result_zero, tensor), "zero ratio is identity" + # positive ratio: should modify + result = sharpen_tensor(tensor.clone(), ratio=0.5) + assert result.shape == tensor.shape, "shape preserved" + assert not torch.isnan(result).any(), "no NaN" + assert not torch.isinf(result).any(), "no Inf" + assert not torch.equal(result, tensor), "sharpening modifies tensor" + return True + + +def test_maximize_tensor(): + """maximize_tensor normalizes tensor range.""" + from modules.processing_correction import maximize_tensor + tensor = torch.randn(4, 64, 64) * 0.5 + # boundary 1.0: identity + result_id = maximize_tensor(tensor.clone(), boundary=1.0) + assert torch.allclose(result_id, tensor), "boundary 1.0 is identity" + # boundary 2.0: should expand range + result = maximize_tensor(tensor.clone(), boundary=2.0) + assert result.abs().max() > tensor.abs().max(), "boundary 2.0 expands range" + assert result.shape == tensor.shape, "shape preserved" + assert not torch.isnan(result).any(), "no NaN" + # boundary 0.5: should compress range + result_small = maximize_tensor(tensor.clone(), boundary=0.5) + assert result_small.abs().max() < tensor.abs().max() + 0.1, "boundary 0.5 compresses" + return True + + +# ============================================================ +# Latent Corrections: correction() pipeline with mock p object +# ============================================================ + +def _make_mock_p(**overrides): + """Create a mock processing object with default hdr params.""" + defaults = { + 'hdr_mode': 0, + 'hdr_brightness': 0.0, + 'hdr_color': 0.0, + 'hdr_sharpen': 0.0, + 'hdr_clamp': False, + 'hdr_boundary': 4.0, + 'hdr_threshold': 0.95, + 'hdr_maximize': False, + 'hdr_max_center': 0.6, + 'hdr_max_boundary': 1.0, + 'hdr_color_picker': '#000000', + 'hdr_tint_ratio': 0.0, + 'correction_total_steps': 20, + 'correction_steps_mid': 10, + 'correction_steps_late': 4, + 'extra_generation_params': {}, + } + defaults.update(overrides) + return SimpleNamespace(**defaults) + + +def test_correction_noop(): + """correction() with all-zero params is near-identity.""" + from modules.processing_correction import correction + p = _make_mock_p() + latent = torch.randn(4, 64, 64) + for step in [0, 5, 10, 15, 19]: + result = correction(p, 500, latent.clone(), step=step) + assert torch.allclose(result, latent, atol=1e-5), f"step {step}: no-op correction should be identity" + return True + + +def test_correction_early_clamp(): + """correction() applies soft_clamp in early steps when hdr_clamp=True.""" + from modules.processing_correction import correction + p = _make_mock_p(hdr_clamp=True, hdr_threshold=0.8, hdr_boundary=4.0) + latent = torch.randn(4, 64, 64) + latent[0, 0, 0] = 10.0 # outlier + # step 0 of 20 = progress 0.0 (early) + result = correction(p, 999, latent.clone(), step=0) + assert result[0, 0, 0] < 10.0, "outlier should be clamped" + assert "Latent clamp" in p.extra_generation_params, "clamp recorded in params" + return True + + +def test_correction_mid_color(): + """correction() applies color centering in mid steps.""" + from modules.processing_correction import correction + p = _make_mock_p(hdr_color=0.5) + latent = torch.randn(4, 64, 64) + 1.0 # offset channels + original_ch_means = [latent[c].mean().item() for c in range(1, 4)] + # step 6 of 20 = progress 0.3 (mid range) + result = correction(p, 700, latent.clone(), step=6) + new_ch_means = [result[c].mean().item() for c in range(1, 4)] + # at least some channels should have their mean reduced (centered) + centered_count = sum(1 for o, n in zip(original_ch_means, new_ch_means) if abs(n) < abs(o)) + assert centered_count > 0, "at least one color channel should be more centered" + assert "Latent color" in p.extra_generation_params, "color recorded in params" + return True + + +def test_correction_late_brightness(): + """correction() applies brightness offset in late steps.""" + from modules.processing_correction import correction + p = _make_mock_p(hdr_brightness=2.0) + latent = torch.randn(4, 64, 64) + original_mean = latent[0].mean().item() + # step 17 of 20 = progress 0.85 (late) + result = correction(p, 100, latent.clone(), step=17) + new_mean = result[0].mean().item() + assert new_mean != original_mean, "brightness should shift channel 0 mean" + assert "Latent brightness" in p.extra_generation_params, "brightness recorded in params" + return True + + +def test_correction_sharpen(): + """correction() applies sharpening in sharpen range.""" + from modules.processing_correction import correction + p = _make_mock_p(hdr_sharpen=1.0) + latent = torch.randn(4, 64, 64) + # step 15 of 20 = progress 0.75 (sharpen range) + result = correction(p, 200, latent.clone(), step=15) + assert not torch.equal(result, latent), "sharpening should modify latent" + assert "Latent sharpen" in p.extra_generation_params, "sharpen recorded in params" + return True + + +def test_correction_maximize(): + """correction() applies maximize/normalize in very late steps.""" + from modules.processing_correction import correction + p = _make_mock_p(hdr_maximize=True, hdr_max_center=0.6, hdr_max_boundary=2.0) + latent = torch.randn(4, 64, 64) * 0.5 + # step 19 of 20 = progress 0.95 (very late) + result = correction(p, 10, latent.clone(), step=19) + assert result.abs().max() > latent.abs().max(), "maximize should expand range" + assert "Latent max" in p.extra_generation_params, "maximize recorded in params" + return True + + +def test_correction_multichannel(): + """correction() uses multi-channel path for >4 channel latents.""" + from modules.processing_correction import correction + p = _make_mock_p(hdr_brightness=2.0, hdr_color=0.5) + # 16-channel latent (e.g. Flux 2) + latent = torch.randn(16, 64, 64) + # mid step: color centering on all channels + p.extra_generation_params = {} + result_mid = correction(p, 700, latent.clone(), step=6) + assert result_mid.shape == latent.shape, "multi-channel shape preserved" + assert not torch.isnan(result_mid).any(), "no NaN" + # late step: brightness via multiplicative scaling + p.extra_generation_params = {} + result_late = correction(p, 100, latent.clone(), step=17) + assert result_late.shape == latent.shape, "multi-channel shape preserved" + assert not torch.isnan(result_late).any(), "no NaN" + assert "Latent brightness" in p.extra_generation_params, "brightness recorded" + return True + + +def test_correction_shape_preservation(): + """correction() preserves shape and dtype for various latent sizes.""" + from modules.processing_correction import correction + p = _make_mock_p(hdr_clamp=True, hdr_color=0.3, hdr_brightness=1.0, hdr_sharpen=0.5) + shapes = [(4, 64, 64), (4, 32, 32), (4, 128, 128), (8, 64, 64), (16, 32, 32)] + for shape in shapes: + for step in [0, 6, 15, 19]: + p.extra_generation_params = {} + latent = torch.randn(shape) + result = correction(p, 500, latent.clone(), step=step) + assert result.shape == latent.shape, f"shape {shape} step {step}: shape mismatch" + assert result.dtype == latent.dtype, f"shape {shape} step {step}: dtype mismatch" + assert not torch.isnan(result).any(), f"shape {shape} step {step}: NaN" + assert not torch.isinf(result).any(), f"shape {shape} step {step}: Inf" + return True + + +def test_correction_step_ranges(): + """correction() applies different operations at different progress points.""" + from modules.processing_correction import correction + p = _make_mock_p( + hdr_clamp=True, hdr_color=0.5, hdr_brightness=1.0, + hdr_sharpen=0.5, hdr_maximize=True, hdr_max_boundary=2.0, + ) + latent = torch.randn(4, 64, 64) + latent[0, 0, 0] = 10.0 # outlier for clamp testing + expected_params_per_range = { + 0: ["Latent clamp"], # early: progress 0.0 + 6: ["Latent color"], # mid: progress 0.3 + 15: ["Latent sharpen"], # sharpen: progress 0.75 + 17: ["Latent brightness"], # late: progress 0.85 + 19: ["Latent max"], # very late: progress 0.95 + } + for step, expected_keys in expected_params_per_range.items(): + p.extra_generation_params = {} + correction(p, 500, latent.clone(), step=step) + for key in expected_keys: + assert key in p.extra_generation_params, f"step {step}: expected '{key}' in params, got {list(p.extra_generation_params.keys())}" + return True + + +# ============================================================ +# Test runner +# ============================================================ + +def run_test(fn): + name = fn.__name__ + try: + result = fn() + if result is None: + log.warning(f' SKIP: {name} (dependency not available)') + return + record(True, name) + except AssertionError as e: + record(False, name, str(e)) + except Exception as e: + record(False, name, f"exception: {e}") + import traceback + traceback.print_exc() + + +def run_tests(): + t0 = time.time() + + # Grading params (pure Python, no GPU deps) + set_category('grading_params') + log.warning('=== Color Grading: Params & Utilities ===') + for fn in [test_grading_params_defaults, test_grading_is_active, test_grading_float_coercion, + test_hex_to_rgb, test_kelvin_to_rgb]: + run_test(fn) + + # Grading functions (need torch, some need kornia) + set_category('grading_functions') + log.warning('=== Color Grading: Tensor Operations ===') + for fn in [test_apply_vignette, test_apply_grain, test_apply_color_temp, + test_apply_shadows_midtones_highlights, test_grade_image_pipeline, + test_grade_image_edge_cases]: + run_test(fn) + + # Correction primitives (pure torch) + set_category('correction_primitives') + log.warning('=== Latent Corrections: Primitives ===') + for fn in [test_soft_clamp_tensor, test_center_tensor, test_sharpen_tensor, + test_maximize_tensor]: + run_test(fn) + + # Correction pipeline (mock p object) + set_category('correction_pipeline') + log.warning('=== Latent Corrections: Pipeline ===') + for fn in [test_correction_noop, test_correction_early_clamp, test_correction_mid_color, + test_correction_late_brightness, test_correction_sharpen, test_correction_maximize, + test_correction_multichannel, test_correction_shape_preservation, + test_correction_step_ranges]: + run_test(fn) + + t1 = time.time() + + # Summary + log.warning('=== Results ===') + total_passed = 0 + total_failed = 0 + for cat, data in results.items(): + total_passed += data['passed'] + total_failed += data['failed'] + status = 'PASS' if data['failed'] == 0 else 'FAIL' + log.info(f' {cat}: {data["passed"]} passed, {data["failed"]} failed [{status}]') + log.warning(f'Total: {total_passed} passed, {total_failed} failed in {t1 - t0:.2f}s') + if total_failed > 0: + sys.exit(1) + + +if __name__ == "__main__": + run_tests()