From 11d25afa564812e04140934f7a71b005ce3747e5 Mon Sep 17 00:00:00 2001 From: Vito Sansevero Date: Tue, 9 Dec 2025 06:05:04 -0800 Subject: [PATCH 1/2] feat(api): add unload autotag model endpoint --- py/api.py | 55 +++++++++++++++++++++++++++++++++++++++++++++++++++---- 1 file changed, 51 insertions(+), 4 deletions(-) diff --git a/py/api.py b/py/api.py index f177be1..2df1515 100644 --- a/py/api.py +++ b/py/api.py @@ -482,6 +482,10 @@ class PromptManagerAPI: async def apply_autotag_route(request): return await self.apply_autotag(request) + @routes.post("/prompt_manager/autotag/unload") + async def unload_autotag_model_route(request): + return await self.unload_autotag_model(request) + @routes.get("/prompt_manager/scan_output_dir") async def scan_output_dir_route(request): return await self.scan_output_dir(request) @@ -4266,7 +4270,7 @@ class PromptManagerAPI: Get status of available AutoTag models. GET /prompt_manager/autotag/models - Returns model availability, download status, and configuration. + Returns model availability, download status, loaded status, and configuration. """ try: from .autotag import get_autotag_service @@ -4277,7 +4281,9 @@ class PromptManagerAPI: return web.json_response({ 'success': True, 'models': models_status, - 'default_prompt': service.default_prompt + 'default_prompt': service.default_prompt, + 'model_loaded': service.is_model_loaded(), + 'loaded_model_type': service.get_loaded_model_type() }) except Exception as e: @@ -4359,6 +4365,7 @@ class PromptManagerAPI: Query params: model_type: "gguf" or "hf" prompt: custom prompt text + keep_in_memory: "true" or "false" (default true) - keep model loaded after processing Streams SSE progress updates during processing. """ @@ -4370,6 +4377,7 @@ class PromptManagerAPI: model_type = request.query.get('model_type', 'gguf') custom_prompt = request.query.get('prompt', '') skip_tagged = request.query.get('skip_tagged', 'true').lower() == 'true' + keep_in_memory = request.query.get('keep_in_memory', 'true').lower() == 'true' use_gpu = True async def stream_response(): @@ -4489,9 +4497,14 @@ class PromptManagerAPI: errors += 1 processed += 1 - service.unload_model() + # Only unload model if keep_in_memory is False + if not keep_in_memory: + service.unload_model() + model_status = 'Model unloaded' + else: + model_status = 'Model kept in memory' - yield f"data: {json.dumps({'type': 'complete', 'progress': 100, 'processed': processed, 'tagged': tagged, 'skipped': skipped, 'errors': errors, 'status': 'Complete'})}\n\n" + yield f"data: {json.dumps({'type': 'complete', 'progress': 100, 'processed': processed, 'tagged': tagged, 'skipped': skipped, 'errors': errors, 'status': 'Complete', 'model_status': model_status, 'model_loaded': keep_in_memory})}\n\n" except Exception as e: self.logger.error(f"AutoTag error: {e}") @@ -4651,6 +4664,40 @@ class PromptManagerAPI: 'error': str(e) }, status=500) + async def unload_autotag_model(self, request): + """ + Manually unload the AutoTag model from memory. + POST /prompt_manager/autotag/unload + + Used to free up VRAM when the model is no longer needed. + """ + try: + from .autotag import get_autotag_service + + service = get_autotag_service() + + if not service.is_model_loaded(): + return web.json_response({ + 'success': True, + 'message': 'No model was loaded' + }) + + model_type = service.get_loaded_model_type() + service.unload_model() + + return web.json_response({ + 'success': True, + 'message': f'{model_type.upper()} model unloaded successfully', + 'model_loaded': False + }) + + except Exception as e: + self.logger.error(f"Unload autotag model error: {e}") + return web.json_response({ + 'success': False, + 'error': str(e) + }, status=500) + async def scan_output_dir(self, request): """ Scan ComfyUI output directory for images. From 986663cd972e8d4b61f71a3b25782587a1ee607e Mon Sep 17 00:00:00 2001 From: Vito Sansevero Date: Tue, 9 Dec 2025 06:05:14 -0800 Subject: [PATCH 2/2] feat(gallery): add model memory management options --- web/gallery.html | 61 ++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 61 insertions(+) diff --git a/web/gallery.html b/web/gallery.html index 4bc9b71..3f3db40 100644 --- a/web/gallery.html +++ b/web/gallery.html @@ -600,6 +600,25 @@ + +
+ + +
+
@@ -862,6 +881,7 @@ No sentences, no commentary, no captions. Only tags. Keep length short but descr document.getElementById('applyReviewBtn').addEventListener('click', () => this.applyReviewTags()); document.getElementById('cancelAutoTagBtn').addEventListener('click', () => this.cancelAutoTag()); document.getElementById('cancelDownloadBtn').addEventListener('click', () => this.cancelDownload()); + document.getElementById('unloadModelBtn').addEventListener('click', () => this.unloadModel()); // Close AutoTag modals on outside click ['autoTagModal', 'autoTagReviewModal', 'autoTagProgressModal', 'autoTagDownloadModal'].forEach(modalId => { @@ -3579,6 +3599,9 @@ Seed: ${this.currentMetadata.seed || 'Unknown'}`; async checkAutoTagModels() { const ggufStatus = document.getElementById("ggufModelStatus"); const hfStatus = document.getElementById("hfModelStatus"); + const modelLoadedStatus = document.getElementById("modelLoadedStatus"); + const loadedModelType = document.getElementById("loadedModelType"); + const unloadModelBtn = document.getElementById("unloadModelBtn"); ggufStatus.innerHTML = 'Checking...'; hfStatus.innerHTML = 'Checking...'; @@ -3601,6 +3624,16 @@ Seed: ${this.currentMetadata.seed || 'Unknown'}`; } else { hfStatus.innerHTML = ``; } + + // Update model loaded status + if (data.model_loaded && data.loaded_model_type) { + modelLoadedStatus.classList.remove('hidden'); + loadedModelType.textContent = data.loaded_model_type.toUpperCase(); + unloadModelBtn.classList.remove('hidden'); + } else { + modelLoadedStatus.classList.add('hidden'); + unloadModelBtn.classList.add('hidden'); + } } else { ggufStatus.innerHTML = 'Error'; hfStatus.innerHTML = 'Error'; @@ -3670,9 +3703,36 @@ Seed: ${this.currentMetadata.seed || 'Unknown'}`; this.showModal("autoTagModal"); } + async unloadModel() { + const unloadBtn = document.getElementById("unloadModelBtn"); + unloadBtn.disabled = true; + unloadBtn.textContent = 'Unloading...'; + + try { + const response = await fetch('/prompt_manager/autotag/unload', { + method: 'POST' + }); + const data = await response.json(); + + if (data.success) { + this.showNotification(data.message, 'success'); + await this.checkAutoTagModels(); // Refresh status + } else { + this.showNotification(`Failed to unload model: ${data.error}`, 'error'); + } + } catch (error) { + console.error('Error unloading model:', error); + this.showNotification('Failed to unload model', 'error'); + } finally { + unloadBtn.disabled = false; + unloadBtn.textContent = 'Unload Model'; + } + } + async startAutoTag() { const modelType = document.querySelector('input[name="autoTagModel"]:checked').value; const prompt = document.getElementById('autoTagPrompt').value; + const keepInMemory = document.getElementById('keepModelInMemory').checked; // Check if model is downloaded const response = await fetch('/prompt_manager/autotag/models'); @@ -3700,6 +3760,7 @@ Seed: ${this.currentMetadata.seed || 'Unknown'}`; const formData = new URLSearchParams(); formData.append('model_type', modelType); formData.append('prompt', prompt); + formData.append('keep_in_memory', keepInMemory); this.autoTagState.eventSource = new EventSource(`/prompt_manager/autotag/start?${formData.toString()}`);