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 |
|---|---|---|---|
|
torch.Tensor or sequence of int
|
1-D integer indices of length |
required |
|
torch.Tensor or sequence of int
|
1-D integer indices of the same length. Entry |
required |
|
int
|
Sequence length of |
required |
|
int
|
Sequence length of |
required |
|
int
|
Model width of |
required |
|
int
|
Number of attention heads. Each head attends over the
same live pairs with |
required |
|
float
|
Dropout probability applied to the packed attention
weights in training mode only; |
0.0
|
|
bool
|
Whether the four projections learn a bias. |
True
|
|
int or None
|
Key embed width. Must be |
None
|
|
int or None
|
Value embed width, under the same restriction as
|
None
|
|
bool
|
Layout of batched tensors: |
True
|
|
bool
|
If |
False
|
|
str or None
|
Opaque checkpoint identity, typically
|
None
|
|
Generator or None
|
Isolated RNG for Xavier-uniform on the four
projections. |
None
|
|
int or None
|
How many live edges to gather at once in the packed
softmax and mix. |
None
|
Attributes:
| Name | Type | Description |
|---|---|---|
query_features |
int
|
Sequence length of |
key_features |
int
|
Sequence length of |
embed_dim |
int
|
Model width. |
num_heads |
int
|
Number of attention heads. |
head_dim |
int
|
|
nnz |
int
|
Number of live edges, counting any pair added by
|
dropout |
float
|
Dropout probability on packed attention weights. |
batch_first |
bool
|
Layout flag; default |
add_self_loops |
bool
|
Whether missing self-loops were OR-ed at construction. |
source_index |
Tensor
|
Int64 buffer of key / value positions, length |
target_index |
Tensor
|
Int64 buffer of query positions, length |
q_proj, k_proj, v_proj, out_proj |
Linear
|
Four separate |
identity |
str | None
|
The constructor |
chunk_size |
int | None
|
Edge-chunk length, or |
Raises:
| Type | Description |
|---|---|
Kpnn2Error
|
At construction, if the indices are empty, not 1-D
integers, mismatched in length, out of range, or
duplicated as |
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 or None
|
Isolated RNG for the four Xavier draws. |
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.