Merge pull request #81 from ComfyAssets/feature/AutoTag
Feature/auto tag
This commit is contained in:
@@ -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.
|
||||
|
||||
@@ -600,6 +600,25 @@
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<!-- Keep Model in Memory Option -->
|
||||
<div class="mb-6">
|
||||
<label class="flex items-center justify-between p-4 bg-gray-700 rounded-lg cursor-pointer hover:bg-gray-650 border border-gray-600">
|
||||
<div class="flex items-center">
|
||||
<input type="checkbox" id="keepModelInMemory" checked class="w-4 h-4 text-purple-600 bg-gray-700 border-gray-600 rounded focus:ring-purple-500">
|
||||
<div class="ml-3">
|
||||
<span class="text-gray-100 font-medium">Keep model in memory</span>
|
||||
<span class="text-gray-400 text-sm ml-2">(Faster processing, uses ~4-16GB VRAM)</span>
|
||||
</div>
|
||||
</div>
|
||||
<button id="unloadModelBtn" class="px-3 py-1 bg-red-600 hover:bg-red-700 text-white text-xs rounded transition-colors hidden">
|
||||
Unload Model
|
||||
</button>
|
||||
</label>
|
||||
<p id="modelLoadedStatus" class="text-gray-500 text-xs mt-2 hidden">
|
||||
<span class="text-green-400">●</span> Model currently loaded: <span id="loadedModelType">none</span>
|
||||
</p>
|
||||
</div>
|
||||
|
||||
<!-- Prompt Template -->
|
||||
<div class="mb-6">
|
||||
<label class="block text-sm font-medium text-gray-400 mb-2">Tag Generation Prompt</label>
|
||||
@@ -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 = '<span class="text-gray-400 text-sm">Checking...</span>';
|
||||
hfStatus.innerHTML = '<span class="text-gray-400 text-sm">Checking...</span>';
|
||||
@@ -3601,6 +3624,16 @@ Seed: ${this.currentMetadata.seed || 'Unknown'}`;
|
||||
} else {
|
||||
hfStatus.innerHTML = `<button onclick="window.gallery.downloadModel('hf')" class="px-3 py-1 bg-blue-600 hover:bg-blue-700 text-white text-xs rounded transition-colors">Download</button>`;
|
||||
}
|
||||
|
||||
// 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 = '<span class="text-red-400 text-sm">Error</span>';
|
||||
hfStatus.innerHTML = '<span class="text-red-400 text-sm">Error</span>';
|
||||
@@ -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()}`);
|
||||
|
||||
|
||||
Reference in New Issue
Block a user