From 85cd222793de4226a2075b5a61cb637be45fee12 Mon Sep 17 00:00:00 2001 From: CalamitousFelicitousness Date: Tue, 2 Dec 2025 21:48:09 +0000 Subject: [PATCH] fix(vqa): sort CLiP analysis results and add text output Improvements to the OpenCLIP interrogation: - Sort all ranking dicts by similarity score (descending) - Add format_category() helper for text formatting - Add formatted text output for CLIP labels textbox - Return additional text update in analyze_image() --- modules/interrogate/openclip.py | 27 ++++++++++++++++++++++----- 1 file changed, 22 insertions(+), 5 deletions(-) diff --git a/modules/interrogate/openclip.py b/modules/interrogate/openclip.py index 508a4d897..2350dc440 100644 --- a/modules/interrogate/openclip.py +++ b/modules/interrogate/openclip.py @@ -204,15 +204,32 @@ def analyze_image(image, clip_model, blip_model): top_movements = ci.movements.rank(image_features, 5) top_trendings = ci.trendings.rank(image_features, 5) top_flavors = ci.flavors.rank(image_features, 5) - medium_ranks = dict(zip(top_mediums, ci.similarities(image_features, top_mediums))) - artist_ranks = dict(zip(top_artists, ci.similarities(image_features, top_artists))) - movement_ranks = dict(zip(top_movements, ci.similarities(image_features, top_movements))) - trending_ranks = dict(zip(top_trendings, ci.similarities(image_features, top_trendings))) - flavor_ranks = dict(zip(top_flavors, ci.similarities(image_features, top_flavors))) + medium_ranks = dict(sorted(zip(top_mediums, ci.similarities(image_features, top_mediums)), key=lambda x: x[1], reverse=True)) + artist_ranks = dict(sorted(zip(top_artists, ci.similarities(image_features, top_artists)), key=lambda x: x[1], reverse=True)) + movement_ranks = dict(sorted(zip(top_movements, ci.similarities(image_features, top_movements)), key=lambda x: x[1], reverse=True)) + trending_ranks = dict(sorted(zip(top_trendings, ci.similarities(image_features, top_trendings)), key=lambda x: x[1], reverse=True)) + flavor_ranks = dict(sorted(zip(top_flavors, ci.similarities(image_features, top_flavors)), key=lambda x: x[1], reverse=True)) + + # Format labels as text + def format_category(name, ranks): + lines = [f"{name}:"] + for item, score in ranks.items(): + lines.append(f" • {item} - {score*100:.1f}%") + return '\n'.join(lines) + + formatted_text = '\n\n'.join([ + format_category("Medium", medium_ranks), + format_category("Artist", artist_ranks), + format_category("Movement", movement_ranks), + format_category("Trending", trending_ranks), + format_category("Flavor", flavor_ranks), + ]) + return [ gr.update(value=medium_ranks, visible=True), gr.update(value=artist_ranks, visible=True), gr.update(value=movement_ranks, visible=True), gr.update(value=trending_ranks, visible=True), gr.update(value=flavor_ranks, visible=True), + gr.update(value=formatted_text, visible=True), # New text output for the textbox ]