boundlab.diff.zono3.DiffSoftmaxPruning#

class boundlab.diff.zono3.DiffSoftmaxPruning[source]#

Bases: OpHandler

Differential handler for boundlab::SoftmaxPruning over the last axis.

Network 1 keeps the full softmax; network 2 sees the mask applied to both the numerator and every denominator term, i.e. h(sᵢ)·exp(dᵢ) / Σⱼ h(sⱼ)·exp(dⱼ).

Methods

__init__

condition

Whether this handler applies to these operands (default: always).

handle

Transform the operands; sub-operations go through interp.<op> so the enclosing domain's handlers apply to them too.

override_handler

A copy of this handler that takes precedence over other when both are ready for the same call.

op: str = 'softmax_pruning'#
handle(interp, scores, data, *, dim=-1, **params)[source]#

Transform the operands; sub-operations go through interp.<op> so the enclosing domain’s handlers apply to them too.

condition(*args, **kwargs)#

Whether this handler applies to these operands (default: always).

override_handler(other)#

A copy of this handler that takes precedence over other when both are ready for the same call.

overrides: list[type[OpHandler]] = []#