PyTorch datasets with pre-computed model logits for efficient research
Project description
Dataset with Logits
A PyTorch package for loading computer vision datasets paired with pre-computed model logits. Perfect for knowledge distillation, model analysis, and efficient research workflows.
🚀 Quick Start
pip install dataset-with-logits
import torchvision.transforms as transforms
from dataset_with_logits import ImageNet
# Define transforms
transform = transforms.Compose([
transforms.Resize(256),
transforms.CenterCrop(224),
transforms.ToTensor(),
])
# Create dataset (auto-downloads predictions)
dataset = ImageNet(
root='/path/to/imagenet/val',
model='resnet18',
transform=transform,
auto_download=True
)
# Use with DataLoader
from torch.utils.data import DataLoader
loader = DataLoader(dataset, batch_size=32, shuffle=True)
for images, labels, logits in loader:
# images: [batch_size, 3, 224, 224]
# labels: [batch_size] - ground truth
# logits: [batch_size, 1000] - model predictions
break
📊 Available Models
ImageNet-1K
resnet18- ResNet-18 (11.7M parameters)resnet50- ResNet-50 (25.6M parameters)resnet152- ResNet-152 (60.2M parameters)vit_l_16- Vision Transformer Large (304M parameters)mobilenet_v3_small- MobileNet V3 Small (2.5M parameters)mobilenet_v3_large- MobileNet V3 Large (5.5M parameters)
More models and datasets coming soon!
🎯 Use Cases
Knowledge Distillation
import torch.nn.functional as F
def knowledge_distillation_loss(student_logits, teacher_logits, labels, temperature=3.0):
student_soft = F.log_softmax(student_logits / temperature, dim=1)
teacher_soft = F.softmax(teacher_logits / temperature, dim=1)
return F.kl_div(student_soft, teacher_soft, reduction='batchmean')
# In your training loop
for images, labels, teacher_logits in dataloader:
student_logits = student_model(images)
loss = knowledge_distillation_loss(student_logits, teacher_logits, labels)
Model Analysis
from dataset_with_logits import ImageNet
# Compare different models
models = ['resnet18', 'resnet152', 'vit_l_16']
datasets = {}
for model in models:
datasets[model] = ImageNet(root=imagenet_path, model=model)
# Analyze prediction differences, calibration, etc.
🔧 Advanced Usage
List Available Models
from dataset_with_logits import list_available_models
models = list_available_models()
print(models)
# {'imagenet1k': {'resnet18': 'ResNet-18 (11.7M parameters)', ...}}
Custom Cache Directory
dataset = ImageNet(
root='/path/to/imagenet',
model='resnet18',
cache_dir='/custom/cache/dir',
auto_download=True
)
Version Control
dataset = ImageNet(
root='/path/to/imagenet',
model='resnet18',
version='v0.1.0', # Specific version
auto_download=True
)
📁 File Format
Prediction files are CSV format with:
id: Image filename (no extension)label: Ground truth class indexlogits: Semicolon-separated model outputs
Example:
id,label,logits
ILSVRC2012_val_00000001,65,-2.3;1.7;0.2;...;0.8
ILSVRC2012_val_00000002,970,0.1;-1.2;3.4;...;-0.5
🌐 Data Source
Prediction files are automatically downloaded from the predictions-on-datasets repository using Git LFS. Files are cached locally after first download.
🔍 Examples
See the examples/ directory for:
- Basic usage
- Knowledge distillation
- Model comparison
- Advanced workflows
📦 Installation
From PyPI (Recommended)
pip install dataset-with-logits
From Source
git clone https://github.com/ViGeng/predictions-on-datasets.git
cd predictions-on-datasets/dataset_with_logits
pip install -e .
🤝 Contributing
Contributions are welcome! See the main repository for contribution guidelines.
📄 License
MIT License - see LICENSE file for details.
Project details
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 dataset_with_logits-0.2.5.tar.gz.
File metadata
- Download URL: dataset_with_logits-0.2.5.tar.gz
- Upload date:
- Size: 14.9 kB
- Tags: Source
- Uploaded using Trusted Publishing? No
- Uploaded via: twine/6.1.0 CPython/3.9.23
File hashes
| Algorithm | Hash digest | |
|---|---|---|
| SHA256 |
eda24bac1b859e6f0c1c777c284a4a0c2dcc010ad8a28f2ca05717aa4e7aa668
|
|
| MD5 |
9c166d7a4ff45fd80f7af7db624a3f2f
|
|
| BLAKE2b-256 |
68d0d3ffce7439c2791810fa96a95ca60501208f477353ec1422265050cdf1be
|
File details
Details for the file dataset_with_logits-0.2.5-py3-none-any.whl.
File metadata
- Download URL: dataset_with_logits-0.2.5-py3-none-any.whl
- Upload date:
- Size: 13.3 kB
- Tags: Python 3
- Uploaded using Trusted Publishing? No
- Uploaded via: twine/6.1.0 CPython/3.9.23
File hashes
| Algorithm | Hash digest | |
|---|---|---|
| SHA256 |
513b04e7f220d13d56ef09fae7ad3f1fcba9f4dbf921ac993e56b82793970d6f
|
|
| MD5 |
c53288c4f1ad24e836bd9d18c788fc4a
|
|
| BLAKE2b-256 |
ae91996c31573e38e9a8db9cf4ff5b05665ee9e93e9ccc673d82f5103f1989b0
|