Skip to main content

High-performance sigmoid attention kernels using Triton

Project description

Triton Sigmoid Attention

Python 3.11+ PyTorch 2.6+ License: MIT Code style: black

Fused sigmoid attention kernels for NVIDIA GPUs built with Triton.

Overview

Two implementations:

  • Dense - Same-length sequences (fastest)
  • Padded - Variable-length with padding masks

Features torch.compile support and causal masking. Requires Ampere architecture or newer.

Installation

pip install triton-sigmoid
Development install from source
git clone https://github.com/MSDLLCpapers/triton-sigmoid.git
cd triton-sigmoid
uv sync --extra dev
source .venv/bin/activate

Quick Start

Dense Attention (Same-Length Sequences)

import torch
from triton_sigmoid import sigmoid_attention

batch, seq_len, n_heads, head_dim = 2, 1024, 8, 64
q = torch.randn(batch, seq_len, n_heads, head_dim, device='cuda', dtype=torch.float16)
k = torch.randn(batch, seq_len, n_heads, head_dim, device='cuda', dtype=torch.float16)
v = torch.randn(batch, seq_len, n_heads, head_dim, device='cuda', dtype=torch.float16)

output = sigmoid_attention(q, k, v, is_causal=False)
output_causal = sigmoid_attention(q, k, v, is_causal=True)

Padded Attention (Variable-Length Sequences)

import torch
from triton_sigmoid import sigmoid_attention_padded

batch, seq_len, n_heads, head_dim = 2, 1024, 8, 64
q = torch.randn(batch, seq_len, n_heads, head_dim, device='cuda', dtype=torch.float16)
k = torch.randn(batch, seq_len, n_heads, head_dim, device='cuda', dtype=torch.float16)
v = torch.randn(batch, seq_len, n_heads, head_dim, device='cuda', dtype=torch.float16)

seq_lens_k = torch.tensor([800, 950], device='cuda', dtype=torch.int32)
seq_lens_q = torch.tensor([800, 950], device='cuda', dtype=torch.int32)

output = sigmoid_attention_padded(q, k, v, seq_lens_k=seq_lens_k, seq_lens_q=seq_lens_q)

compiled_fn = torch.compile(sigmoid_attention_padded)
output = compiled_fn(q, k, v, seq_lens_k=seq_lens_k, seq_lens_q=seq_lens_q, is_causal=True)

Performance

TFLOPS Comparison

See benchmarks/README.md for details.

Documentation

License

MIT License - see LICENSE file for details.

Contributing

Contributions are welcome! Please see CONTRIBUTING.md for guidelines.

Citation

If you use this work in your research, please cite our paper:

@misc{sadashivaiah2026bettermodelsfastertraining,
      title={Better Models, Faster Training: Sigmoid Attention for single-cell Foundation Models}, 
      author={Vijay Sadashivaiah and Georgios Dasoulas and Judith Mueller and Soumya Ghosh},
      year={2026},
      eprint={2604.27124},
      archivePrefix={arXiv},
      primaryClass={cs.LG},
      url={https://arxiv.org/abs/2604.27124}, 
}

Acknowledgments

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

triton_sigmoid-0.1.1.tar.gz (30.6 kB view details)

Uploaded Source

Built Distribution

If you're not sure about the file name format, learn more about wheel file names.

triton_sigmoid-0.1.1-py3-none-any.whl (21.8 kB view details)

Uploaded Python 3

File details

Details for the file triton_sigmoid-0.1.1.tar.gz.

File metadata

  • Download URL: triton_sigmoid-0.1.1.tar.gz
  • Upload date:
  • Size: 30.6 kB
  • Tags: Source
  • Uploaded using Trusted Publishing? Yes
  • Uploaded via: twine/6.1.0 CPython/3.13.12

File hashes

Hashes for triton_sigmoid-0.1.1.tar.gz
Algorithm Hash digest
SHA256 8c0f95e41f0bb8f8c4e30b1af361e27ce439398e3ee003895cfac6ecc611c4bf
MD5 99c19273ff38b3d4657e6611bcac8b15
BLAKE2b-256 9e9f7d1f9aae36d78801657e53aafa8881cf00298ab340b52897d13547f7a5ba

See more details on using hashes here.

Provenance

The following attestation bundles were made for triton_sigmoid-0.1.1.tar.gz:

Publisher: publish.yml on MSDLLCpapers/triton-sigmoid

Attestations: Values shown here reflect the state when the release was signed and may no longer be current.

File details

Details for the file triton_sigmoid-0.1.1-py3-none-any.whl.

File metadata

  • Download URL: triton_sigmoid-0.1.1-py3-none-any.whl
  • Upload date:
  • Size: 21.8 kB
  • Tags: Python 3
  • Uploaded using Trusted Publishing? Yes
  • Uploaded via: twine/6.1.0 CPython/3.13.12

File hashes

Hashes for triton_sigmoid-0.1.1-py3-none-any.whl
Algorithm Hash digest
SHA256 310e6d5b6564e2495d0a96141917a22bae73de25c57b6ff3f28e3760a1459224
MD5 2bc8d883523591cc0b0145d92b0802a2
BLAKE2b-256 1f9cd3c4d23fbfca12aed60159918b0c59d325287589ff20e329cace7f3e9a0e

See more details on using hashes here.

Provenance

The following attestation bundles were made for triton_sigmoid-0.1.1-py3-none-any.whl:

Publisher: publish.yml on MSDLLCpapers/triton-sigmoid

Attestations: Values shown here reflect the state when the release was signed and may no longer be current.

Supported by

AWS Cloud computing and Security Sponsor Datadog Monitoring Depot Continuous Integration Fastly CDN Google Download Analytics Pingdom Monitoring Sentry Error logging StatusPage Status page