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()
This commit is contained in:
CalamitousFelicitousness
2025-12-02 21:48:09 +00:00
parent eb832a4850
commit 85cd222793
+22 -5
View File
@@ -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
]