v2.1.3 - Count what an attention layer computes instead of reporting it as free (#157)
🌟 Summary
🚀 THOP 2.1.3 delivers more accurate PyTorch operation profiling, especially for attention-based models, while improving pooling, recurrent layers, custom rules, and profiler reliability.
📊 Key Changes
-
🧠 Accurate
nn.MultiheadAttentioncounting (PR #157)
Attention is no longer reported as free. THOP now counts its projection layers, attention matrix operations, and softmax processing, including support for:- Batched and unbatched inputs
- Cross-attention
- Different key and value dimensions
- Optional bias and zero-attention tokens
- Calls with or without returned attention weights
- Positional, mixed, and keyword arguments
-
📈 More realistic attention-model estimates
Because attention cost does not scale linearly with image area, stride-based shortcut estimation is disabled for models containing attention. These models now use the exact target-size profiling path. For example, measured RT-DETR estimates increased by roughly 1.5–1.8%, reflecting previously uncounted work. -
🧮 Improved average-pooling operation counts (PR #150)
Fixed and adaptive average pooling now account for actual window sizes, padding, boundary behavior,ceil_mode, andcount_include_pad. This avoids assuming every output uses a full pooling window. -
🧱 Better support for unbatched and recurrent inputs (PR #150)
Convolution and RNN/GRU/LSTM profilers now correctly interpret channel, batch, and sequence dimensions for both batched and unbatched layouts. -
⚡ Expanded zero-operation layer coverage (PR #152)
Data movement and selection layers such asFlatten,Unflatten,Identity,PixelUnshuffle, additional padding layers, and fractional max pooling are explicitly registered as zero-MAC operations. This reduces unnecessary missing-rule warnings and improves profiling speed. -
🔥 Complete softmax-family support (PR #153)
LogSoftmax,Softmin, andSoftmax2dare now counted, including their additional elementwise work and correct handling of empty normalization dimensions. -
🛡️ Safer handling of caller-owned
total_opsvalues (PR #151)
Existingtotal_opsattributes and buffers are temporarily preserved during profiling and restored afterward, preventing THOP from overwriting model state. -
🌳 Profiler results now reflect the modules that actually ran (PR #154)
Both profiler entry points consistently count the forward pass that executed, even when a model replaces or detaches submodules during inference. Training modes are also restored more reliably. -
🧩 More flexible custom operation rules (PRs #155–#156)
Custom counting rules can now be callable objects orfunctools.partialinstances, and they receive arguments supplied through keyword, positional, or mixed calls. -
📦 Version updated to
2.1.3
🎯 Purpose & Impact
- ✅ More trustworthy FLOPs/MACs reports, particularly for transformer and attention-based architectures such as RT-DETR.
- 🎯 Better model comparison and deployment planning, since previously omitted attention and softmax costs are now included.
- 🧪 More reliable profiling across PyTorch versions and input layouts, including older, unbatched, recurrent, and keyword-heavy workloads.
- ⚡ Faster profiling for layers that genuinely perform no arithmetic, with fewer unnecessary hooks and warnings.
- 🔒 Safer integration into existing projects, as profiling no longer permanently overwrites caller-defined
total_opsdata or model training states.
What's Changed
- Read the axes each layer owns and charge the average pool window by @raimbekovm in #150
- Preserve a caller-owned total_ops instead of colliding with it by @raimbekovm in #151
- Register the zero-op layers that reached no counting rule by @raimbekovm in #152
- Count the whole softmax family and stop dividing by an empty axis by @raimbekovm in #153
- Answer for the tree the forward pass ran, not the one it left behind by @raimbekovm in #154
- Accept any callable as a custom_ops counting rule by @raimbekovm in #155
- Hand a counting rule the arguments its module was called with by @raimbekovm in #156
- Count what an attention layer computes instead of reporting it as free by @raimbekovm in #157
Full Changelog: v2.1.2...v2.1.3