Loss Functions

Registry

.. py:function:: get_loss_function(loss_type: str, **kwargs: ~typing.Any) -> ~collections.abc.Callable[[~torch.Tensor, ~torch.Tensor], ~torch.Tensor]

module:

silverspoon_kd.losses

Get a loss function by type name.

param loss_type:

Name of the loss function. Canonical names: mse, normalized_mse, cosine, smooth_l1, kl_divergence, jsd, logit_lens_kl, contrastive, angular_magnitude, mahalanobis_mse, mahalanobis_cosine, relkd_distance, relkd_angle, relkd_distance_angle. Aliases kl_div, mahal_mse, mahal_cosine, relkd_da are also accepted.

param **kwargs:

Additional arguments passed to the loss function factory

returns:

A callable loss function that takes (student_output, teacher_output) tensors

raises ValueError:

If loss_type is not recognized

.. rubric:: Example

loss_fn = get_loss_function("mse")
loss_fn = get_loss_function("kl_divergence", temperature=3.0)
loss_fn = get_loss_function("mahalanobis_mse", weight_matrix=W)

.. py:data:: LOSS_REGISTRY

module:

silverspoon_kd.losses

value:

{‘angular_magnitude’: , ‘contrastive’: , ‘cosine’: , ‘jsd’: , ‘kl_div’: , ‘kl_divergence’: , ‘logit_lens_kl’: , ‘mahal_cosine’: , ‘mahal_mse’: , ‘mahalanobis_cosine’: , ‘mahalanobis_mse’: , ‘mse’: , ‘normalized_mse’: , ‘relkd_angle’: , ‘relkd_da’: , ‘relkd_distance’: , ‘relkd_distance_angle’: , ‘smooth_l1’: }

dict() -> new empty dictionary dict(mapping) -> new dictionary initialized from a mapping object’s (key, value) pairs dict(iterable) -> new dictionary initialized as if via: d = {} for k, v in iterable: d[k] = v dict(**kwargs) -> new dictionary initialized with the name=value pairs in the keyword argument list. For example: dict(one=1, two=2)

Built-in losses

.. py:function:: mse_loss(*, reduction: str = ‘mean’) -> ~collections.abc.Callable[[~torch.Tensor, ~torch.Tensor], ~torch.Tensor]

module:

silverspoon_kd.losses

Mean Squared Error loss.

param reduction:

Reduction applied to the per-element loss ('mean', 'sum', or 'none'). Default: 'mean'.

returns:

A callable (student, teacher) -> scalar loss function.

.. py:function:: normalized_mse_loss(*, eps: float = 1e-05) -> ~collections.abc.Callable[[~torch.Tensor, ~torch.Tensor], ~torch.Tensor]

module:

silverspoon_kd.losses

Normalized MSE loss (z-score both tensors, then MSE).

Matches relational structure (which tokens differ from the mean) without being distracted by absolute activation magnitudes. Addresses the main flaw of plain MSE for quantized or scaled representations.

param eps:

Numerical-stability term added to the standard deviation denominator. Default: 1e-5.

returns:

A callable (student, teacher) -> scalar loss function.

.. py:function:: cosine_loss(*, dim: int = -1) -> ~collections.abc.Callable[[~torch.Tensor, ~torch.Tensor], ~torch.Tensor]

module:

silverspoon_kd.losses

Cosine similarity loss (1 - cosine_similarity).

param dim:

Dimension along which cosine similarity is computed. Default: -1.

returns:

A callable (student, teacher) -> scalar loss function.

.. py:function:: smooth_l1_loss(*, beta: float = 1.0, reduction: str = ‘mean’) -> ~collections.abc.Callable[[~torch.Tensor, ~torch.Tensor], ~torch.Tensor]

module:

silverspoon_kd.losses

Smooth L1 (Huber) loss.

param beta:

Transition point between L2 (|x| < beta) and L1 behaviour. Default: 1.0.

param reduction:

Reduction applied to the per-element loss ('mean', 'sum', or 'none'). Default: 'mean'.

returns:

A callable (student, teacher) -> scalar loss function.

.. py:function:: kl_divergence_loss(*, temperature: float = 1.0, chunk_size: int = 0) -> ~collections.abc.Callable[[~torch.Tensor, ~torch.Tensor], ~torch.Tensor]

module:

silverspoon_kd.losses

KL Divergence loss for logit distillation.

Supports optional chunked computation for large vocabularies to reduce peak memory usage.

param temperature:

Temperature for softmax scaling. Default: 1.0.

param chunk_size:

Number of tokens per chunk for memory-efficient computation. Set to 0 to disable chunking. Default: 0.

returns:

A callable (student, teacher) -> scalar loss function.

.. py:function:: logit_lens_kl_loss(*, output_head: ~torch.nn.modules.module.Module | None = None, temperature: float = 1.0) -> ~collections.abc.Callable[[~torch.Tensor, ~torch.Tensor], ~torch.Tensor]

module:

silverspoon_kd.losses

Logit-lens KL loss (project through output head, then compare distributions).

Inspired by the FDD trajectory loss (Gong et al., ACL 2025). Projects both student and teacher hidden states through a shared output head and compares the resulting token distributions via KL divergence. This measures functional equivalence — whether the layer outputs mean the same thing for prediction — rather than raw numerical similarity.

param output_head:

Required. The output head to project through (e.g. an LM head for language models).

param temperature:

Softmax temperature. Default: 1.0.

returns:

A callable (student, teacher) -> scalar loss function.

.. py:function:: jsd_loss(*, temperature: float = 1.0, beta: float = 0.5, chunk_size: int = 0) -> ~collections.abc.Callable[[~torch.Tensor, ~torch.Tensor], ~torch.Tensor]

module:

silverspoon_kd.losses

Jensen-Shannon Divergence loss for logit distillation.

JSD(P||Q) = beta * KL(P||M) + (1 - beta) * KL(Q||M) where M = beta * P + (1 - beta) * Q.

Supports optional chunked computation for large vocabularies.

param temperature:

Temperature for softmax scaling. Default: 1.0.

param beta:

Interpolation weight strictly in (0, 1). 0.5 is symmetric JSD; values near 0/1 bias toward one direction of the KL. Default: 0.5.

param chunk_size:

Number of tokens per chunk. Set to 0 to disable. Default: 0.

returns:

A callable (student, teacher) -> scalar loss function.

.. py:class:: ContrastiveDistillationLoss(temperature: float = 0.07, student_dim: int | None = None, teacher_dim: int | None = None)

module:

silverspoon_kd.losses

canonical:

silverspoon_kd.losses.contrastive.ContrastiveDistillationLoss

Bases: :py:class:~torch.nn.modules.module.Module

InfoNCE-based contrastive loss for knowledge distillation.

Trains the student to produce representations that are similar to the corresponding teacher representations (positive pairs) while being dissimilar to other samples in the batch (negative pairs).

.. rubric:: References

  • CRD: Contrastive Representation Distillation (ICLR 2020)

  • CoDIR: Contrastive Distillation on Intermediate Representations (EMNLP 2020)

.. note::

For sequence inputs (3D tensors), this implementation flattens all tokens into independent samples. If you need sequence-level representations, apply pooling (e.g., mean or CLS token) before passing to this loss.

param temperature:

Temperature for softmax scaling. Lower values make the distribution sharper. Default: 0.07.

param student_dim:

Dimensionality of student features. If None, inferred on first forward pass.

param teacher_dim:

Dimensionality of teacher features. If None, inferred on first forward pass.

.. rubric:: Example

# Simplest usage (dimensions inferred automatically):
loss_fn = ContrastiveDistillationLoss()

# With explicit dimensions (creates projector upfront if they differ):
loss_fn = ContrastiveDistillationLoss(student_dim=256, teacher_dim=512)

.. py:method:: ContrastiveDistillationLoss.init(temperature: float = 0.07, student_dim: int | None = None, teacher_dim: int | None = None)

module:

silverspoon_kd.losses

Initialize internal Module state, shared by both nn.Module and ScriptModule.

.. py:method:: ContrastiveDistillationLoss.forward(student_features: ~torch.Tensor, teacher_features: ~torch.Tensor) -> ~torch.Tensor

module:

silverspoon_kd.losses

Compute contrastive distillation loss.

param student_features:

Student representations of shape [B, D_s] or [B, L, D_s].

param teacher_features:

Teacher representations of shape [B, D_t] or [B, L, D_t].

returns:

Scalar loss tensor.

.. py:function:: contrastive_loss(*, temperature: float = 0.07, student_dim: int | None = None, teacher_dim: int | None = None) -> ~silverspoon_kd.losses.contrastive.ContrastiveDistillationLoss

module:

silverspoon_kd.losses

Contrastive distillation loss (InfoNCE-based).

param temperature:

Softmax temperature. Default: 0.07.

param student_dim:

Student feature dimension. If None, inferred on first forward pass.

param teacher_dim:

Teacher feature dimension. If None, inferred on first forward pass.

returns:

A ContrastiveDistillationLoss instance.

.. py:function:: angular_magnitude_loss(*, alpha: float = 1.0, beta: float = 1.0, dim: int = -1) -> ~collections.abc.Callable[[~torch.Tensor, ~torch.Tensor], ~torch.Tensor]

module:

silverspoon_kd.losses

Decomposed angular + magnitude loss.

Separately matches direction (cosine) and scale (norm difference), giving explicit control over both components that plain MSE conflates.

param alpha:

Weight for the angular (cosine) term. Default: 1.0.

param beta:

Weight for the magnitude term. Default: 1.0.

param dim:

Dimension along which cosine similarity is computed. Default: -1.

returns:

A callable (student, teacher) -> scalar loss function.

.. py:function:: mahal_mse_loss(*, weight_matrix: ~torch.Tensor | None = None, metric_matrix: ~torch.Tensor | None = None, pre_norm: ~torch.nn.modules.module.Module | None = None) -> ~collections.abc.Callable[[~torch.Tensor, ~torch.Tensor], ~torch.Tensor]

module:

silverspoon_kd.losses

MSE loss under Mahalanobis metric M = W^T W.

Computes (s - t)^T M (s - t) averaged over all positions, which is equivalent to MSE in logit space but at O(D²) cost instead of O(V·D).

Architecture-agnostic: operates on raw hidden states. The metric matrix is supplied as an argument, not derived from model internals.

param weight_matrix:

Output-projection weights (V, D). M = W^T W is computed once at construction. Mutually exclusive with metric_matrix.

param metric_matrix:

Precomputed PSD metric matrix (D, D). Mutually exclusive with weight_matrix.

param pre_norm:

Optional normalization module applied to both student and teacher before the metric (e.g. the model’s final RMSNorm / LayerNorm).

returns:

A callable (student, teacher) -> scalar loss function.

.. py:function:: mahal_cosine_loss(*, weight_matrix: ~torch.Tensor | None = None, metric_matrix: ~torch.Tensor | None = None, pre_norm: ~torch.nn.modules.module.Module | None = None, eps: float = 1e-08) -> ~collections.abc.Callable[[~torch.Tensor, ~torch.Tensor], ~torch.Tensor]

module:

silverspoon_kd.losses

Cosine loss under Mahalanobis metric M = W^T W.

Computes 1 - cos_M(s, t) where cos_M(s, t) = s^T M t / sqrt(s^T M s · t^T M t), averaged over all positions. Equivalent to cosine similarity in logit space at O(D²) instead of O(V·D).

Architecture-agnostic: operates on raw hidden states.

param weight_matrix:

Output-projection weights (V, D). M = W^T W is computed once at construction. Mutually exclusive with metric_matrix.

param metric_matrix:

Precomputed PSD metric matrix (D, D). Mutually exclusive with weight_matrix.

param pre_norm:

Optional normalization module applied to both student and teacher before the metric.

param eps:

Numerical-stability term for the denominator. Default: 1e-8.

returns:

A callable (student, teacher) -> scalar loss function.

.. py:function:: relkd_distance_loss() -> ~collections.abc.Callable[[~torch.Tensor, ~torch.Tensor], ~torch.Tensor]

module:

silverspoon_kd.losses

RelKD distance-wise distillation loss.

Penalizes differences in pairwise Euclidean distance structure between teacher and student representations.

.. warning:: Relational losses require batch_size >= 2 because they operate on pairwise relations between examples. With batch_size=1 the pairwise distance matrix is trivially zero and the loss is always 0 regardless of student quality — effectively a no-op step.

returns:

A callable (student, teacher) -> scalar loss function.

.. py:function:: relkd_angle_loss() -> ~collections.abc.Callable[[~torch.Tensor, ~torch.Tensor], ~torch.Tensor]

module:

silverspoon_kd.losses

RelKD angle-wise distillation loss.

Penalizes differences in the angle formed by triplets of examples between teacher and student representations.

.. warning:: Relational losses require batch_size >= 2 because they operate on pairwise relations between examples. With batch_size=1 the angle tensor is trivially zero and the loss is always 0.

returns:

A callable (student, teacher) -> scalar loss function.

.. py:function:: relkd_da_loss(*, dist_weight: float = 1.0, angle_weight: float = 2.0) -> ~collections.abc.Callable[[~torch.Tensor, ~torch.Tensor], ~torch.Tensor]

module:

silverspoon_kd.losses

Combined RelKD distance + angle loss (RelKD-DA).

Default weights follow the paper: dist_weight=1, angle_weight=2.

param dist_weight:

Weight for the distance term. Default: 1.0.

param angle_weight:

Weight for the angle term. Default: 2.0.

returns:

A callable (student, teacher) -> scalar loss function.

Liger Kernel

.. py:class:: FusedLinearKLDivLoss(*args, **kwargs)

module:

silverspoon_kd.losses

canonical:

silverspoon_kd.losses.liger.FusedLinearKLDivLoss

Bases: :py:class:~torch.nn.modules.module.Module

Placeholder — the real class requires liger-kernel.

.. py:method:: FusedLinearKLDivLoss.init(*args, **kwargs)

module:

silverspoon_kd.losses

Initialize internal Module state, shared by both nn.Module and ScriptModule.

.. py:class:: LigerFusedLinearJSDLoss(*args, **kwargs)

module:

silverspoon_kd.losses

canonical:

silverspoon_kd.losses.liger.LigerFusedLinearJSDLoss

Bases: :py:class:~torch.nn.modules.module.Module

Placeholder — the real class requires liger-kernel.

.. py:method:: LigerFusedLinearJSDLoss.init(*args, **kwargs)

module:

silverspoon_kd.losses

Initialize internal Module state, shared by both nn.Module and ScriptModule.