A PyTorch device manager for automatic hardware detection and memory optimization
Project description
Torch Device Manager
A lightweight PyTorch utility for automatic hardware detection and memory optimization across different devices (CPU, CUDA, MPS).
Features
- 🔍 Automatic Device Detection: Detects the best available hardware (CUDA, Apple Silicon MPS, or CPU)
- 🧠 Memory Optimization: Automatically adjusts batch sizes and gradient accumulation based on available memory
- ⚡ Mixed Precision Support: Optional automatic mixed precision with gradient scaling
- 📊 Memory Monitoring: Real-time memory usage tracking and logging
- 🛡️ Fallback Protection: Graceful fallback to CPU when requested devices aren't available
Installation
pip install torch-device-manager
Quick Start
from torch_device_manager import DeviceManager
import torch
# Initialize device manager (auto-detects best device)
device_manager = DeviceManager(device="auto", mixed_precision=True)
# Get the torch device
device = device_manager.get_device()
# Move your model to the optimal device
model = YourModel().to(device)
# Optimize batch size based on available memory
optimal_batch_size, gradient_steps = device_manager.optimize_for_memory(
model=model,
batch_size=32
)
print(f"Using device: {device}")
print(f"Optimized batch size: {optimal_batch_size}")
print(f"Gradient accumulation steps: {gradient_steps}")
Usage in Training Scripts
Basic Integration
import torch
import torch.nn as nn
from torch_device_manager import DeviceManager
def train_model():
# Initialize device manager
device_manager = DeviceManager(device="auto", mixed_precision=True)
device = device_manager.get_device()
# Setup model
model = YourModel().to(device)
optimizer = torch.optim.Adam(model.parameters())
# Optimize memory usage
batch_size, gradient_steps = device_manager.optimize_for_memory(model, 32)
# Training loop
for epoch in range(num_epochs):
for batch_idx, (data, target) in enumerate(dataloader):
data, target = data.to(device), target.to(device)
# Use mixed precision if available
if device_manager.mixed_precision and device_manager.scaler:
with torch.cuda.amp.autocast():
output = model(data)
loss = criterion(output, target)
device_manager.scaler.scale(loss).backward()
if (batch_idx + 1) % gradient_steps == 0:
device_manager.scaler.step(optimizer)
device_manager.scaler.update()
optimizer.zero_grad()
else:
output = model(data)
loss = criterion(output, target)
loss.backward()
if (batch_idx + 1) % gradient_steps == 0:
optimizer.step()
optimizer.zero_grad()
# Log memory usage
device_manager.log_memory_usage()
Advanced Usage
from torch_device_manager import DeviceManager
# Force specific device
device_manager = DeviceManager(device="cuda", mixed_precision=False)
# Check memory info
memory_info = device_manager.get_memory_info()
print(f"Available memory: {memory_info}")
# Manual memory optimization
if memory_info.get("free_gb", 0) < 2.0:
print("Low memory detected, reducing batch size")
batch_size = 4
API Reference
DeviceManager
Constructor
device(str, default="auto"): Device to use ("auto", "cuda", "mps", "cpu")mixed_precision(bool, default=True): Enable mixed precision training
Methods
get_device(): Returns torch.device objectget_memory_info(): Returns memory information dictlog_memory_usage(): Logs current memory usageoptimize_for_memory(model, batch_size): Returns optimized (batch_size, gradient_steps)
Device Support
- CUDA: Full support with memory optimization and mixed precision
- Apple Silicon (MPS): Basic support with conservative memory settings
- CPU: Fallback support with optimized batch sizes
Requirements
- Python >= 3.8
- PyTorch >= 2.1.1
License
MIT License
Contributing
Contributions are welcome!
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 torch_device_manager-0.1.0.tar.gz.
File metadata
- Download URL: torch_device_manager-0.1.0.tar.gz
- Upload date:
- Size: 7.7 kB
- Tags: Source
- Uploaded using Trusted Publishing? Yes
- Uploaded via: twine/6.1.0 CPython/3.12.9
File hashes
| Algorithm | Hash digest | |
|---|---|---|
| SHA256 |
31bbdec42054c410e12f800ad80b7efd44b48d6caf6a05b77d33f45cacb02960
|
|
| MD5 |
fee22521d42bf36f4d028b2065f8b0a8
|
|
| BLAKE2b-256 |
9172d12f8a63d9d98c9ed509b721e5e4c1362324ab9ef85bf9245952e1c18623
|
Provenance
The following attestation bundles were made for torch_device_manager-0.1.0.tar.gz:
Publisher:
publish.yml on TempCoder82/torch-device-manager
-
Statement:
-
Statement type:
https://in-toto.io/Statement/v1 -
Predicate type:
https://docs.pypi.org/attestations/publish/v1 -
Subject name:
torch_device_manager-0.1.0.tar.gz -
Subject digest:
31bbdec42054c410e12f800ad80b7efd44b48d6caf6a05b77d33f45cacb02960 - Sigstore transparency entry: 293938826
- Sigstore integration time:
-
Permalink:
TempCoder82/torch-device-manager@b8382f7c38747a0fcb06756170d00351b310944f -
Branch / Tag:
refs/tags/v0.1.0 - Owner: https://github.com/TempCoder82
-
Access:
public
-
Token Issuer:
https://token.actions.githubusercontent.com -
Runner Environment:
github-hosted -
Publication workflow:
publish.yml@b8382f7c38747a0fcb06756170d00351b310944f -
Trigger Event:
release
-
Statement type:
File details
Details for the file torch_device_manager-0.1.0-py3-none-any.whl.
File metadata
- Download URL: torch_device_manager-0.1.0-py3-none-any.whl
- Upload date:
- Size: 5.8 kB
- Tags: Python 3
- Uploaded using Trusted Publishing? Yes
- Uploaded via: twine/6.1.0 CPython/3.12.9
File hashes
| Algorithm | Hash digest | |
|---|---|---|
| SHA256 |
6fc555669c43b7e81ee7a26ca8897628f8d45bb8b4aa9bf481c26c500c820858
|
|
| MD5 |
4e71811b275b29f08e6d91b041649674
|
|
| BLAKE2b-256 |
667cf33aa97897703c0ef9d234642ec8310d4f3214603f0671cc5b148a27da52
|
Provenance
The following attestation bundles were made for torch_device_manager-0.1.0-py3-none-any.whl:
Publisher:
publish.yml on TempCoder82/torch-device-manager
-
Statement:
-
Statement type:
https://in-toto.io/Statement/v1 -
Predicate type:
https://docs.pypi.org/attestations/publish/v1 -
Subject name:
torch_device_manager-0.1.0-py3-none-any.whl -
Subject digest:
6fc555669c43b7e81ee7a26ca8897628f8d45bb8b4aa9bf481c26c500c820858 - Sigstore transparency entry: 293938860
- Sigstore integration time:
-
Permalink:
TempCoder82/torch-device-manager@b8382f7c38747a0fcb06756170d00351b310944f -
Branch / Tag:
refs/tags/v0.1.0 - Owner: https://github.com/TempCoder82
-
Access:
public
-
Token Issuer:
https://token.actions.githubusercontent.com -
Runner Environment:
github-hosted -
Publication workflow:
publish.yml@b8382f7c38747a0fcb06756170d00351b310944f -
Trigger Event:
release
-
Statement type: