Skip to content

PackedMultiheadAttention

PackedMultiheadAttention(
    source_index: object,
    target_index: object,
    query_features: int,
    key_features: int,
    embed_dim: int,
    num_heads: int,
    dropout: float = 0.0,
    bias: bool = True,
    kdim: int | None = None,
    vdim: int | None = None,
    batch_first: bool = True,
    add_self_loops: bool = False,
    *,
    identity: str | None = None,
    generator: Generator | None = None,
    chunk_size: int | None = None
)

Bases: Module

Multi-head attention restricted to the live edges of a named edgelist.

Prior knowledge decides which node may attend to which: the query at an edge's target sees only the keys at its sources. Reach for it instead of masking nn.MultiheadAttention, whose mask costs an (n, n) matrix. Build it from parse_adjacency indices and wrap it in your own encoder. Defaults differ from MHA: batch_first=True and need_weights=False. True returns packed per-edge weights, not a dense (L, S) matrix.

Parameters:

Name Type Description Default

source_index

torch.Tensor or sequence of int

1-D integer indices of length nnz >= 1. Entry i is the key / value position of live edge i, and must satisfy 0 <= source_index < key_features. Copied to an int64 buffer, so later writes to the argument do not reach this layer.

required

target_index

torch.Tensor or sequence of int

1-D integer indices of the same length. Entry i is the query position of live edge i, and must satisfy 0 <= target_index < query_features. The two arrays are paired position by position and must not repeat a (source, target) pair.

required

query_features

int

Sequence length of query, that is, how many query positions the packed indices address. Positive int.

required

key_features

int

Sequence length of key / value. Positive int. Equal to query_features for a self-attention graph over one node set; smaller or larger for a bipartite query/key map.

required

embed_dim

int

Model width of query / key / value and of the output. Positive int, divisible by num_heads.

required

num_heads

int

Number of attention heads. Each head attends over the same live pairs with embed_dim // num_heads channels.

required

dropout

float

Dropout probability applied to the packed attention weights in training mode only; 0.0 disables it. Must be >= 0; integer 0 is accepted, bool and negatives are rejected.

0.0

bias

bool

Whether the four projections learn a bias. False makes them pure linear maps.

True

kdim

int or None

Key embed width. Must be None or equal to embed_dim; kept for nn.MultiheadAttention call-site parity, not to support a differing width.

None

vdim

int or None

Value embed width, under the same restriction as kdim.

None

batch_first

bool

Layout of batched tensors: (..., seq, embed_dim) when True (kpnn2 sample-major), (seq, batch, embed_dim) when False (the nn.MultiheadAttention layout). Unbatched 2-D (seq, embed) ignores this flag.

True

add_self_loops

bool

If True, OR any missing (i, i) pair into the module buffers, so every query keeps its own token as a key. Existing self-loops are kept, not duplicated, and the caller's index objects are not mutated. Requires query_features == key_features.

False

identity

str or None

Opaque checkpoint identity, typically spec.fingerprint. Stored in state_dict next to index_digest as a 1-D CPU uint8 tensor of the UTF-8 bytes. load_state_dict raises Kpnn2Error when a present identity does not match this layer, and does not load the weights. A missing identity is not an error, even with strict=True. None means this layer does not claim an identity.

None

generator

Generator or None

Isolated RNG for Xavier-uniform on the four projections. None uses the default torch generator and keeps today's path: nn.Linear kaiming-init then Xavier on that global stream. When set, those nn.Linear constructors do not advance the global stream. Not stored on the module; pass it again to reset_parameters to replay. Do not pass a seed integer.

None

chunk_size

int or None

How many live edges to gather at once in the packed softmax and mix. None gathers all live pairs at once. A positive int uses chunked live pairs: softmax is still over every live key of a query, but pair gathers are slices of that many edges and those gathers are rematerialized in backward. bool and values other than None or a positive int raise Kpnn2Error. Not stored on the module state_dict or in index_digest.

None

Attributes:

Name Type Description
query_features int

Sequence length of query.

key_features int

Sequence length of key / value.

embed_dim int

Model width.

num_heads int

Number of attention heads.

head_dim int

embed_dim // num_heads.

nnz int

Number of live edges, counting any pair added by add_self_loops.

dropout float

Dropout probability on packed attention weights.

batch_first bool

Layout flag; default True.

add_self_loops bool

Whether missing self-loops were OR-ed at construction.

source_index Tensor

Int64 buffer of key / value positions, length nnz. Treat it as read-only: like any PyTorch buffer it can be written to, and doing so rewires the layer without reinitializing it. Rebuild from the edgelist instead.

target_index Tensor

Int64 buffer of query positions, length nnz, read-only in the same sense.

q_proj, k_proj, v_proj, out_proj Linear

Four separate embed_dim -> embed_dim projections, not a fused in_proj_weight.

identity str | None

The constructor identity, or None.

chunk_size int | None

Edge-chunk length, or None for all live pairs.

Raises:

Type Description
Kpnn2Error

At construction, if the indices are empty, not 1-D integers, mismatched in length, out of range, or duplicated as (source, target) pairs; if the sizes are not positive ints; if embed_dim is not divisible by num_heads; if dropout is a bool or negative; if kdim / vdim are neither None nor embed_dim; if add_self_loops is set when query_features != key_features; if identity is neither a str nor None; if generator is neither a torch.Generator nor None; or if chunk_size is neither None nor a positive int. From load_state_dict, when the checkpoint carries an index digest or identity that does not match this layer, in which case the weights are not loaded. forward documents its own rejected arguments.

See Also

PackedLinear : Same packed pairs when the update is one trainable scalar per edge rather than a contraction. MaskedLinear : Dense (out, in) weight; the default layer whenever that square fits. AdjacencySpec : Supplies source_index / target_index; this layer takes those tuples, not the spec object. torch.nn.MultiheadAttention : Dense equivalent, and the reference for call shape, heads, and training.

Notes

The mix is a softmax over each query's live keys only, formed edge by edge, so no (n, n) or (L, S) score matrix is allocated and torch.sparse is not imported; cost scales with nnz, not with query_features * key_features. chunk_size=None (the default) gathers all live pairs at once. A positive chunk_size gathers chunked live pairs and rematerializes those (..., chunk, heads, head_dim) tensors in backward so training does not save a full (..., nnz, heads, head_dim) gather. Packed softmax (..., nnz, heads) is still stored. A query with no live keys (and none added by add_self_loops) mixes to zeros rather than NaN, then still goes through out_proj. Input nodes of an AdjacencySpec are exactly that case.

forward documents the accepted layouts, packed need_weights, attn_mask / is_causal, and key_padding_mask.

Index buffers stay integer after .half() / bfloat16 / .double(); projection weights follow the module floating dtype like nn.Linear. torch.autocast is unsupported: forward disables it and casts query, key, and value to the parameter dtype. Cast the module with .to(dtype=...) (or .half() / .double()) instead. state_dict adds index_digest, a 1-D CPU uint8 tensor of length 32: the SHA-256 of the packed indices plus query_features, key_features, embed_dim, and num_heads, so a reshape or a different head split cannot collide. It is not a registered buffer, and a missing digest is not an error, even with strict=True. identity is stored the same way when the constructor was given one: a rename that leaves the packed index pattern unchanged is caught when callers pass spec.fingerprint. chunk_size is not part of state_dict or index_digest. copy.deepcopy works.

Examples:

An input feeding a two-node feedback core plus one output, attended as one four-token sequence:

>>> import pandas as pd
>>> import torch
>>> import kpnn2
>>> edgelist = pd.DataFrame(
...     {
...         "source": ["x", "a", "b", "a"],
...         "target": ["a", "b", "a", "y"],
...     }
... )
>>> spec = kpnn2.parse_adjacency(edgelist)
>>> n = spec.state_dim
>>> attn = kpnn2.PackedMultiheadAttention(
...     spec.source_index,
...     spec.target_index,
...     n,
...     n,
...     embed_dim=8,
...     num_heads=2,
...     add_self_loops=True,
...     identity=spec.fingerprint,
... )
>>> attn.nnz, attn.head_dim
(8, 4)
>>> tokens = torch.randn(2, n, 8)
>>> output, weights = attn(tokens, tokens, tokens)
>>> tuple(output.shape)
(2, 4, 8)
>>> weights is None
True
>>> _, packed = attn(
...     tokens,
...     tokens,
...     tokens,
...     need_weights=True,
... )
>>> tuple(packed.shape)
(2, 8)

Methods:

reset_parameters

reset_parameters(
    generator: Generator | None = None,
) -> None

Xavier-uniform projections, like nn.MultiheadAttention.

Parameters:

Name Type Description Default

generator

Generator or None

Isolated RNG for the four Xavier draws. None uses the default torch generator. Not stored on the module. Biases, when present, are zeroed and do not consume the generator.

None

forward

forward(
    query: Tensor,
    key: Tensor,
    value: Tensor,
    key_padding_mask: Tensor | None = None,
    need_weights: bool = False,
    attn_mask: Tensor | None = None,
    average_attn_weights: bool = True,
    is_causal: bool = False,
) -> tuple[torch.Tensor, torch.Tensor | None]

Packed multi-head attention; drop-in call shape vs MHA.

query last dim is embed_dim; sequence length is query_features. key and value last dim is embed_dim; sequence length is key_features.

With batch_first=True (the default), tensors are (..., seq, embed_dim), including unbatched 2-D (seq, embed) and batched 3-D (batch, seq, embed). With batch_first=False, unbatched 2-D is still (seq, embed); batched 3-D is (seq, batch, embed), transposed to batch-first for the packed kernel and transposed back.

Always returns a 2-tuple. need_weights defaults to False, and then the second entry is None. If need_weights is True, the second entry is the packed per-edge softmax, aligned with source_index / target_index, not MHA's dense (L, S) map. With average_attn_weights=True (the default) the head axis is averaged, shape (..., nnz). With False, shape (..., nnz, num_heads). Batch layout follows the output (including batch_first). This path does not allocate (L, S).

attn_mask must be None. is_causal must be False. key_padding_mask is None or a boolean mask: True means ignore that key. Shape (S,) unbatched or (N, S) batched. Applied in packed space; does not allocate (n, n). Copied onto the scores' device; stays boolean. Float padding masks raise Kpnn2Error. After padding, a query with no remaining keys stays zeros, not NaN.