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 |
|---|---|---|---|
|
Tensor
|
Connectivity of shape |
required |
|
bool
|
If |
True
|
|
str or None
|
Opaque checkpoint identity, typically
|
None
|
|
Module or None
|
Optional per-entry map on |
None
|
|
Generator or None
|
Isolated RNG for |
None
|
Attributes:
| Name | Type | Description |
|---|---|---|
in_features |
int
|
Number of input columns, |
out_features |
int
|
Number of output columns, |
weight |
Parameter
|
The trainable tensor, shape
|
mask |
Tensor
|
Float32 buffer, same shape as the constructor |
constraint |
Module or None
|
The constructor |
bias |
Parameter | None
|
Trainable bias, or |
identity |
str | None
|
The constructor |
Raises:
| Type | Description |
|---|---|
Kpnn2Error
|
If |
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_dictkeys areweight, optionalbias,mask_digest, andidentitywhen the constructor was given one;maskstays out of it.mask_digestis a 1-D CPUuint8tensor 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 withstrict=True. The digest catches same-shape rewiring. A rename that leaves the 0/1 pattern unchanged is caught byidentitywhen callers passspec.fingerprint. A kpnn2 0.1 checkpoint, which stored the trainable tensor asparametrizations.weight.original, still loads; its masked-out entries are set to 0 on load, which does not change the output.- Param-group filters match
weightthe same way on this layer,PackedLinear, andnn.Linear. copy.deepcopyand pickling (torch.save(model)) work.torch.nn.utils.pruneworks onweight; 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 |
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 or None
|
Isolated RNG for these draws. |
None
|
Raises:
| Type | Description |
|---|---|
Kpnn2Error
|
If |
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.