Generate attention heatmaps for WD Tagger models - visualize which parts of an image contribute to each tag
Project description
WD Tag Heatmap
Generate attention heatmaps for WD Tagger models - visualize which parts of an image contribute to each tag prediction.
Features
- 🎯 Generate individual attention heatmaps for each detected tag
- 🖼️ Create grid visualizations of all tag heatmaps
- 📊 Export detailed tag information with confidence scores
- 🚀 Support for multiple WD Tagger v3 models
- 💻 Easy to use Python API and CLI
- 🔧 Batch processing support
Installation
pip install wd-tag-heatmap
For development:
git clone https://github.com/yourusername/wd-tag-heatmap
cd wd-tag-heatmap
pip install -e ".[dev]"
Quick Start
Python API
from wd_tag_heatmap import generate_tag_heatmaps
# Generate heatmaps for a single image
heatmaps, grid, labels = generate_tag_heatmaps(
"anime_girl.jpg",
output_dir="results",
threshold=0.35
)
print(f"Found {len(heatmaps)} tags")
print(f"Caption: {labels.caption}")
Command Line
# Process single image
wd-tag-heatmap image.jpg
# With custom settings
wd-tag-heatmap image.jpg --output results --threshold 0.5 --model SmilingWolf/wd-vit-large-tagger-v3
# Batch process
wd-tag-heatmap *.jpg --output batch_results
Usage Examples
Basic Usage
from wd_tag_heatmap import generate_tag_heatmaps
# Simple usage with defaults
heatmaps, grid, labels = generate_tag_heatmaps("my_image.png")
Custom Configuration
from wd_tag_heatmap import generate_tag_heatmaps, AVAILABLE_MODELS
# Use a different model and threshold
heatmaps, grid, labels = generate_tag_heatmaps(
"my_image.jpg",
output_dir="custom_output",
model_repo="SmilingWolf/wd-convnext-tagger-v3",
threshold=0.5,
save_individual_tags=True,
save_grid=True,
save_tags_info=True
)
# Access results programmatically
for heatmap in heatmaps[:5]:
print(f"{heatmap.label}: {heatmap.score:.3f}")
Batch Processing
from wd_tag_heatmap import batch_process_images
# Process multiple images
image_paths = ["img1.jpg", "img2.png", "img3.jpg"]
batch_process_images(
image_paths,
output_dir="batch_results",
threshold=0.35
)
Working with Results
# Get results without saving files
heatmaps, grid, labels = generate_tag_heatmaps(
"image.jpg",
save_individual_tags=False,
save_grid=False,
save_tags_info=False
)
# Access tag information
print(f"Rating tags: {labels.rating}")
print(f"Character tags: {labels.character}")
print(f"General tags: {list(labels.general.keys())[:10]}")
# Work with heatmap images
for heatmap in heatmaps[:3]:
# heatmap.image is a PIL Image
heatmap.image.show()
Available Models
SmilingWolf/wd-convnext-tagger-v3SmilingWolf/wd-swinv2-tagger-v3SmilingWolf/wd-vit-tagger-v3(default)SmilingWolf/wd-vit-large-tagger-v3SmilingWolf/wd-eva02-large-tagger-v3
Output Structure
output_dir/
├── image_name_heatmap_grid.png # Grid of all heatmaps
├── image_name_tags_info.txt # Detailed tag information
└── individual_tags/ # Individual heatmap images
├── image_name_000_1girl_0.996.png
├── image_name_001_solo_0.975.png
└── ...
Requirements
- Python 3.8+
- PyTorch 2.0+
- CUDA GPU recommended for faster processing
API Reference
generate_tag_heatmaps()
Main function to generate heatmaps for a single image.
Parameters:
image_path(str): Path to input imageoutput_dir(str): Directory for outputs (default: "heatmap_outputs")model_repo(str): Model repository (default: "SmilingWolf/wd-vit-tagger-v3")threshold(float): Confidence threshold (default: 0.35)save_individual_tags(bool): Save individual heatmaps (default: True)save_grid(bool): Save grid image (default: True)save_tags_info(bool): Save tag information (default: True)
Returns:
heatmaps(List[Heatmap]): List of heatmap objectsgrid(Image): Grid visualizationlabels(ImageLabels): Tag information
batch_process_images()
Process multiple images in batch.
Parameters:
image_paths(List[str]): List of image paths- Other parameters same as
generate_tag_heatmaps()
Contributing
Contributions are welcome! Please feel free to submit a Pull Request.
License
This project is licensed under the MIT License - see the LICENSE file for details.
Acknowledgments
- WD Tagger models by SmilingWolf
- Built with timm
Citation
If you use this tool in your research, please cite:
@software{wd-tag-heatmap,
title = {WD Tag Heatmap: Attention Visualization for Anime Image Tagging},
author = {Your Name},
year = {2024},
url = {https://github.com/yourusername/wd-tag-heatmap}
}
Project details
Release history Release notifications | RSS feed
Download files
Download the file for your platform. If you're not sure which to choose, learn more about installing packages.
Source Distribution
Built Distribution
Filter files by name, interpreter, ABI, and platform.
If you're not sure about the file name format, learn more about wheel file names.
Copy a direct link to the current filters
File details
Details for the file wd_tag_heatmap-0.1.0.tar.gz.
File metadata
- Download URL: wd_tag_heatmap-0.1.0.tar.gz
- Upload date:
- Size: 14.9 kB
- Tags: Source
- Uploaded using Trusted Publishing? No
- Uploaded via: twine/6.1.0 CPython/3.12.3
File hashes
| Algorithm | Hash digest | |
|---|---|---|
| SHA256 |
f5d468fe596e23786ac8acc76bba5f7e130b2387a2e46bed0a62e8bb53d82656
|
|
| MD5 |
852de35526a250968d698ced09133441
|
|
| BLAKE2b-256 |
ce3ca18eb9a2819b18cd3d3924b0fb7d3e0d24ddac928189d780e41868cce1d7
|
File details
Details for the file wd_tag_heatmap-0.1.0-py3-none-any.whl.
File metadata
- Download URL: wd_tag_heatmap-0.1.0-py3-none-any.whl
- Upload date:
- Size: 14.4 kB
- Tags: Python 3
- Uploaded using Trusted Publishing? No
- Uploaded via: twine/6.1.0 CPython/3.12.3
File hashes
| Algorithm | Hash digest | |
|---|---|---|
| SHA256 |
72d79d7b12eabc9596ad0e00e7ed2f673bb4c23cccd1678138f16ac38f39b2f8
|
|
| MD5 |
6fa506e31062a15661551e49574c585e
|
|
| BLAKE2b-256 |
2a8b9dd6391c3a0d64a058cb0da10338fd8db7aaa92e88408ae927d61149ddf8
|