boundlab.diff.softmax_pruning#

boundlab.diff.softmax_pruning(scores, data, dim=-1)[source]#

Mock softmax pruning: network 1 is softmax(data), network 2 masks it.

Network 2 computes the mask-renormalised softmax h(sⱼ)·exp(dⱼ) / Σₖ h(sₖ)·exp(dₖ). Eagerly this evaluates the pruned network.