Skip to content

v2.1.3 - Count what an attention layer computes instead of reporting it as free (#157)

Choose a tag to compare

@github-actions github-actions released this 29 Jul 18:50
· 18 commits to main since this release
a5dff05

🌟 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.MultiheadAttention counting (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, and count_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 as Flatten, 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, and Softmax2d are now counted, including their additional elementwise work and correct handling of empty normalization dimensions.

  • 🛡️ Safer handling of caller-owned total_ops values (PR #151)
    Existing total_ops attributes 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 or functools.partial instances, 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_ops data 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