Skip to content

MaskedLinear

MaskedLinear(
    mask: Tensor,
    bias: bool = True,
    *,
    identity: str | None = None,
    constraint: Module | None = None,
    generator: Generator | None = None
)

Bases: Module

Affine hop whose connectivity is fixed by an edgelist mask.

One hop of a knowledge-primed network: an nn.Linear-style layer (call layer(x); not a subclass) in which only edges present in the prior-knowledge graph can carry weight, so absent edges need no hand-zeroing after every optimizer step. Build one from spec.hops[i].to_mask() or spec.to_mask(), fed by gather_hop_inputs on a layered hop. The large-n path on the same indices is PackedLinear, which stores its weight the same way: weight is the trainable tensor and effective_weight() is what forward uses. Initialization scales by per-row mask degree.

Parameters:

Name Type Description Default

mask

Tensor

Connectivity of shape (out_features, in_features), usually spec.hops[i].to_mask() or spec.to_mask(). A nonzero entry [j, k] lets input column k reach output row j; a zero blocks it for the life of the layer. Stored as an independent float32 copy, so later writes to the tensor passed in do not reach this layer. Non-finite values are not special-cased: the stored tensor is multiplied with the weight as it is.

required

bias

bool

If True, learn a bias of shape (out_features,). If False, there is no bias.

True

identity

str or None

Opaque checkpoint identity, typically spec.fingerprint. Stored in state_dict next to mask_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

constraint

Module or None

Optional per-entry map on weight, applied in forward before the mask. The module must return a tensor of the same shape as weight. Give it a right_inverse (the torch.nn.utils.parametrize convention) and reset_parameters stores right_inverse(draw), so the effective weight keeps the degree-aware init; without one, weight stores the draw and the map is applied on top of it. nn.Softplus() is the textbook non-negative edge reparametrization but has no right_inverse: every live edge then starts near ln 2 whatever its fan-in (see Notes). PackedLinear takes the same argument. Mixed per-edge signs and frozen slots belong in this module, not in parse columns. A hard freeze is torch.where replacing those slots; a gradient hook that zeroes a slot is not a freeze (AdamW and SGD with momentum still move the stored tensor).

None

generator

Generator or None

Isolated RNG for reset_parameters. None uses the default torch generator, bit-identical to omitting the argument. Not stored on the module; pass it again to reset_parameters to replay. Do not pass a seed integer.

None

Attributes:

Name Type Description
in_features int

Number of input columns, mask.shape[1].

out_features int

Number of output columns, mask.shape[0].

weight Parameter

The trainable tensor, shape (out_features, in_features), stored under the name weight as on nn.Linear. Masked-out entries start at exactly 0; their gradient is always 0, so SGD, momentum, Adam, and weight decay keep them at 0. When constraint is set, this tensor is unconstrained. Read effective_weight() for the map forward uses.

mask Tensor

Float32 buffer, same shape as the constructor mask. Not trained and not saved in state_dict. Stays float32 after .half() / bfloat16 / .double(). Treat it as read-only: like any PyTorch buffer it can be written to, and doing so silently rewires the layer. Rebuild from the edgelist instead.

constraint Module or None

The constructor constraint module, or None.

bias Parameter | None

Trainable bias, or None when constructed with bias=False.

identity str | None

The constructor identity, or None.

Raises:

Type Description
Kpnn2Error

If mask is not a torch.Tensor or is not 2-D; if identity is neither a str nor None; if constraint is neither an nn.Module nor None, or does not preserve the weight shape; if generator is neither a torch.Generator nor None; and from load_state_dict when the checkpoint carries a mask digest or identity that does not match this layer; the weights are then not loaded.

See Also

PackedLinear : One trainable scalar per live edge, for graphs whose dense (out_features, in_features) weight would not fit in RAM. Tied decode is PackedLinear.transpose; on this layer use F.linear(h, layer.effective_weight().T, dec_bias). gather_hop_inputs : Assembles the input tensor of a hop that reads more than one saved layer. torch.nn.Linear : Dense equivalent, and the reference for shapes, bias, calling the module, and training.

Notes

Sizes come from mask.shape; there are no separate size arguments. Forward is Y = F.linear(X, effective_weight(), bias), that is Y = X @ (C(W) ⊙ M).T + b with W the weight parameter, C the constraint (the identity when omitted), and M the mask cast to W's dtype and device, so .half(), bfloat16, and .double() work as on nn.Linear. torch.autocast is unsupported: forward disables it and casts x to the parameter dtype. The multiply is dense on purpose, and X is an ordinary dense activation tensor. Nothing in the forward path is a tensor subclass, so torch.compile(layer, fullgraph=True) traces it.

The mask is always applied last. A map you register_parametrization on weight yourself runs when weight is read, before constraint and before the mask, so it cannot resurrect a blocked entry. constraint is the supported per-entry map for sign-constrained edges.

weight is a plain nn.Parameter, as on nn.Linear and PackedLinear:

  • state_dict keys are weight, optional bias, mask_digest, and identity when the constructor was given one; mask stays out of it. mask_digest is a 1-D CPU uint8 tensor of length 32: the SHA-256 of the live mask's float32 C-contiguous bytes at save time, not a registered buffer. A missing digest is not an error, even with strict=True. The digest catches same-shape rewiring. A rename that leaves the 0/1 pattern unchanged is caught by identity when callers pass spec.fingerprint. A kpnn2 0.1 checkpoint, which stored the trainable tensor as parametrizations.weight.original, still loads; its masked-out entries are set to 0 on load, which does not change the output.
  • Param-group filters match weight the same way on this layer, PackedLinear, and nn.Linear.
  • copy.deepcopy and pickling (torch.save(model)) work.
  • torch.nn.utils.prune works on weight; its mask composes with the connectivity mask.

weight is the unconstrained tensor, so optimizer weight_decay pulls it toward 0. Under softplus that pulls the effective weight toward ln 2, not toward 0. Give a constrained layer weight_decay=0 in its param group and penalize effective_weight() in the loss instead. A right_inverse on the constraint keeps the degree-aware init; a constraint without one can initialize weight itself from init_bound().

reset_parameters uses per-row mask degree as fan_in, not full in_features. Because a hop carries every parent of its target, including skip parents, that per-row degree is the unit's real fan-in. Optional generator isolates those draws from other torch RNG consumers.

Examples:

One hop whose second output is connected only to the second input:

>>> import torch
>>> import kpnn2
>>> mask = torch.tensor(
...     [
...         [1.0, 1.0],
...         [0.0, 1.0],
...     ]
... )
>>> layer = kpnn2.MaskedLinear(
...     mask,
...     bias=False,
... )
>>> layer.in_features, layer.out_features
(2, 2)
>>> y = layer(torch.ones(3, 2))
>>> tuple(y.shape)
(3, 2)

weight is the trainable parameter; a blocked entry starts at 0 there and is 0 in the map forward uses:

>>> [name for name, _ in layer.named_parameters()]
['weight']
>>> bool(layer.weight[1, 0] == 0.0)
True
>>> bool(layer.effective_weight()[1, 0] == 0.0)
True

constraint is applied before the mask, so blocked entries stay zero even though softplus(0) > 0:

>>> layer = kpnn2.MaskedLinear(
...     mask,
...     bias=False,
...     constraint=torch.nn.Softplus(),
... )
>>> bool(layer.effective_weight()[1, 0] == 0.0)
True

Methods:

effective_weight

effective_weight() -> torch.Tensor

Return the (out_features, in_features) map forward uses.

constraint(weight) when constraint is set (else weight), times the mask cast to that tensor's dtype and device. Recomputed on every call and differentiable, so it is the tensor to read, export, or penalize as the layer's live edge weights. PackedLinear has the same method.

Returns:

Type Description
Tensor

The effective weight. Blocked entries are 0.

init_bound

init_bound() -> torch.Tensor

Return the degree-aware init bound of every weight entry.

A live entry of row j gets 1 / sqrt(fan_in) of that row, the bound reset_parameters draws its effective weight from. Blocked entries, and rows with fan_in == 0, get 0. Use it to initialize a constraint that has no right_inverse yourself. PackedLinear has the same method on its packed slots.

Returns:

Type Description
Tensor

Shape (out_features, in_features), dtype and device of weight.

reset_parameters

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

Initialize from per-row mask degree, not full width.

This is the difference from torch.nn.Linear.reset_parameters. For output row j, fan_in is the number of ones in mask[j]. The live entries of that row of weight (and bias[j], if present) are drawn uniformly from [-1 / sqrt(fan_in), 1 / sqrt(fan_in)] (init_bound()). Masked-out entries are 0. If fan_in == 0, the row and bias entry stay 0.

Rows are drawn one at a time, each as a full row that is then masked, so the draws match earlier releases and no full-size temporary is allocated.

The draw is the degree-aware value for the effective weight. Without constraint, or when constraint has no right_inverse, weight stores the draw itself. When constraint defines right_inverse (the torch.nn.utils.parametrize convention), weight stores constraint.right_inverse(draw) on live entries (blocked entries stay 0), so effective_weight() keeps the degree-aware scale. The draws, and so the RNG stream, are the same either way.

Parameters:

Name Type Description Default

generator

Generator or None

Isolated RNG for these draws. None uses the default torch generator. Not stored on the module.

None

Raises:

Type Description
Kpnn2Error

If generator is neither a torch.Generator nor None, or constraint.right_inverse returns a tensor of the wrong shape or a non-finite value on a live entry.

forward

forward(x: Tensor) -> torch.Tensor

F.linear of x with effective_weight().

The constraint (if any) and then the mask are applied to weight here, on every call. The stored mask remains float32 and is cast to the parameter dtype, so .half(), bfloat16, and .double() match nn.Linear. torch.autocast is unsupported: this path disables it and casts x to the parameter dtype.