Mediumword2vec

Skip-gram Negative Sampling Loss

Word2Vec

Medium

Problem

Implement the Skip-gram with Negative Sampling loss for one positive center-context pair and a set of sampled negative words.

\mathcal{L}=-\log\sigma(u_o^{\mathsf T}v_c)-\sum_{i=1}^{K}\log\sigma(-u_i^{\mathsf T}v_c)

Here, v_c is center_vec, is pos_vec, are the rows of neg_vecs, $$ is the number of negatives, and \sigma is the logistic sigmoid. Return the loss as a scalar float64 PyTorch tensor and keep it finite for large finite dot products.

Theory

Skip-gram with Negative Sampling (SGNS) is the training objective introduced by Mikolov et al. (2013) that made word2vec fast enough to train on billions of words. Instead of predicting a full probability distribution over the entire vocabulary for each context word, SGNS reframes training as a set of independent binary classification problems: tell real (center, context) pairs apart from randomly sampled noise pairs.


The Full-Softmax Problem

The original skip-gram model predicts context words from a center word. For a center word c and a context word o, the model uses an input embedding v_c and an output embedding u_o, and defines the probability with a softmax over the whole vocabulary W:

p(o \mid c) = \frac{\exp(u_o^\top v_c)}{\sum_{w=1}^{W} \exp(u_w^\top v_c)}

The denominator sums over every word in the vocabulary. Real vocabularies have 10^5 to 10^7 words, so each training example would require a dot product against every output vector, plus the same cost again during backpropagation. With billions of training tokens this is computationally hopeless.

The paper states the goal plainly: the softmax normalization is "impractical because the cost of computing the gradient is proportional to W". Every alternative in the paper exists to avoid touching all W output vectors per update.


From NCE to Negative Sampling

Negative sampling is a simplification of Noise Contrastive Estimation (NCE, Gutmann and Hyvarinen, 2012; Mnih and Teh, 2012). NCE reduces density estimation to binary classification: train a logistic classifier to separate true data samples from samples drawn from a known noise distribution. NCE preserves enough structure to approximate the softmax probabilities.

word2vec only needs good embeddings, not calibrated probabilities. So the authors drop the parts of NCE that are needed for density estimation and keep only the binary classification core. The result is negative sampling: a cheaper objective tuned specifically for learning representations rather than modeling likelihood.

Concretely, NCE weights each noise sample by the ratio of data and noise densities so that the classifier's output can be converted back into a probability. Negative sampling throws away those weights and the noise normalization, treating every sampled negative as a plain label-0 example. The paper notes that while NCE approximately maximizes the log-probability of the softmax, negative sampling does not, and that this trade is acceptable precisely because the embeddings, not the probabilities, are the product.


The SGNS Objective

For a single observed (center, positive context) pair (c, o), the model draws k negative words n_1, \dots, n_k from a noise distribution and minimizes:

L = -\log\sigma(v_c^\top u_o) - \sum_{i=1}^{k}\log\sigma(-v_c^\top u_{n_i})

where \sigma(x) = 1 / (1 + e^{-x}) is the logistic sigmoid. The terms have a clean interpretation:

Each term is the binary cross-entropy of a logistic classifier. The positive pair has label 1, every negative pair has label 0. The center word uses its input embedding v_c; both the positive and negative context words use their output embeddings u.

A subtle but important point: this loss is for a single observed (center, positive) pair. A skip-gram pass over a sentence generates one such pair for every center word and every context word inside its window, and each of those pairs gets its own fresh set of k negatives. The total training objective sums this per-pair loss over the whole corpus. Because the negatives are resampled per pair, the same noise word can appear as a negative for many different centers, which is fine: the objective only ever asks whether a specific (center, candidate) dot product should be high or low.

Note also that there is no normalization across the vocabulary anywhere in L. Each term depends only on the one dot product it contains. This locality is exactly what makes the gradient cheap: an update touches the input vector of the center, the output vector of the positive, and the k output vectors of the negatives, and nothing else.


Sigmoid as Binary Classification

The sigmoid turns a raw dot product into a probability that a pair is "real". Define D = 1 for a true pair and D = 0 for a noise pair. Then:

p(D = 1 \mid c, w) = \sigma(v_c^\top u_w), \qquad p(D = 0 \mid c, w) = 1 - \sigma(v_c^\top u_w) = \sigma(-v_c^\top u_w)

The identity 1 - \sigma(x) = \sigma(-x) is why the negative term carries the minus sign inside the sigmoid. Maximizing the log-likelihood of the labels (1 for the observed pair, 0 for each sampled negative) gives exactly the loss above. Minimizing $$ is maximizing that log-likelihood.


Why softplus for Stability

Computing -\log\sigma(x) naively means evaluating \sigma(x) first, then taking its log. When x is a large negative number, \sigma(x) underflows to exactly 0 in floating point, and \log(0) = -\infty. The loss then becomes \text{inf} or \text{nan}, and gradients explode. With unbounded embeddings, dot products can easily reach magnitudes of 30 or more, so this is a real failure mode, not a corner case.

The stable rewrite uses the softplus function \text{softplus}(x) = \log(1 + e^x):

-\log\sigma(x) = -\log\frac{1}{1 + e^{-x}} = \log(1 + e^{-x}) = \text{softplus}(-x)

So the positive term becomes \text{softplus}(-v_c^\top u_o) and each negative term becomes \text{softplus}(v_c^\top u_{n_i}):

L = \text{softplus}(-v_c^\top u_o) + \sum_{i=1}^{k}\text{softplus}(v_c^\top u_{n_i})

Library softplus implementations use the identity \text{softplus}(x) = \max(x, 0) + \log(1 + e^{-|x|}), which never overflows or underflows for any finite input. This is the same trick PyTorch uses inside \texttt{F.softplus} and \texttt{F.logsigmoid}.


The Role of k

k is the number of negative samples drawn per positive pair. It directly controls the cost: each update touches k + 1 output vectors instead of all W. The paper reports:

Larger k gives a stronger contrastive signal and more stable gradients, at higher per-step cost. The negatives are sampled from a unigram distribution raised to the power 3/4, which the paper found empirically better than the plain unigram or uniform distribution: it boosts rare words slightly while still favoring frequent ones.


Input vs Output Embeddings

word2vec maintains two separate embedding matrices. The input matrix holds v_w vectors used when a word acts as the center. The output matrix holds u_w vectors used when a word acts as context (positive or negative). The score of a pair is always a dot product between one input vector and one output vector, v_c^\top u_w, never two vectors from the same matrix.

Keeping the matrices separate avoids a self-similarity artifact: a single shared matrix would push a word's vector to have a large dot product with itself, which is undesirable. After training, practitioners usually keep the input matrix as the word vectors, or average the two.


Comparison with Hierarchical Softmax

The same paper proposes hierarchical softmax as the other fast alternative. It arranges the vocabulary as the leaves of a binary Huffman tree and computes p(o \mid c) as a product of sigmoids along the root-to-leaf path, costing O(\log W) per example instead of O(W).

Negative sampling became the default in practice because of its simplicity and strong empirical results, and the same contrastive idea later reappeared throughout representation learning.


Modern Context

SGNS is the prototype for the contrastive learning family that now dominates self-supervised representation learning. The recipe is the same: define a score for pairs, label real co-occurrences as positives, draw negatives from a noise distribution, and train a logistic objective to separate them.

Understanding the SGNS loss is therefore not just historical: the binary-classification-against-noise pattern is a foundational primitive in representation learning.


Gradients and What They Learn

The gradient of the loss makes the contrastive behavior explicit. For the positive pair, the derivative of -\log\sigma(v_c^\top u_o) with respect to the score s_o = v_c^\top u_o is \sigma(s_o) - 1, a value in (-1, 0). For a negative pair, the derivative of -\log\sigma(-v_c^\top u_{n_i}) with respect to s_{n_i} = v_c^\top u_{n_i} is \sigma(s_{n_i}), a value in (0, 1).

By the chain rule, the gradient flowing into the input vector v_c is:

\frac{\partial L}{\partial v_c} = (\sigma(s_o) - 1)\, u_o + \sum_{i=1}^{k} \sigma(s_{n_i})\, u_{n_i}

The interpretation is direct. The term (\sigma(s_o) - 1) is negative, so a gradient step moves v_c toward u_o, pulling the real pair closer. Each \sigma(s_{n_i}) is positive, so the step moves v_c away from the noise vectors u_{n_i}. The magnitude of each push is the model's current error on that pair: a negative that already scores low contributes almost nothing, while a confidently wrong one contributes a near-unit gradient. This automatic focus on hard, informative examples is what makes the objective sample-efficient.


Worked Example (D = 2, k = 1)

Let v_c = [1, 0], u_o = [1, 0], and one negative u_{n} = [-1, 0].

  1. Positive score: v_c^\top u_o = 1 \cdot 1 + 0 \cdot 0 = 1.

  2. Negative score: v_c^\top u_n = 1 \cdot (-1) + 0 \cdot 0 = -1.

  3. Positive loss: \text{softplus}(-1) = \log(1 + e^{-1}) \approx 0.3133.

  4. Negative loss: \text{softplus}(-1) = \log(1 + e^{-1}) \approx 0.3133 (note the negative term uses +v_c^\top u_n = -1 as the softplus input).

  5. Total: L \approx 0.3133 + 0.3133 = 0.6266. The model is doing well here: the real pair has a high score and the noise pair a low score, so both losses are small.

A useful sanity check: if every embedding is the zero vector, all dot products are 0, \sigma(0) = 0.5, and the loss equals (k+1)\log 2 \approx 0.6931(k+1).


Pitfalls


Examples

Example 1

Input
center_vec = [0.5,-0.2], pos_vec = [0.3,0.4], neg_vecs = [[-0.1,0.2],[0.4,-0.3]]
Output
2.139492
Explanation
The loss combines one positive softplus term with the sum of the two negative terms.

Example 2

Input
center_vec = [0.2,0.1,-0.4], pos_vec = [0.5,-0.2,0.3], neg_vecs = [[0.1,0.1,0.1]]
Output
1.401507

Example 3

Input
center_vec = [0.1,-0.2,0.3,-0.4], pos_vec = [0.2,0.2,-0.1,0.5], neg_vecs = [[0.1,0,-0.2,0.3],[-0.4,0.1,0.2,-0.1],[0,0.3,-0.3,0.2]]
Output
2.735787

Hints

  1. Use torch.dot for the positive pair and a matrix-vector product for all negatives.
  2. The stable loss is F.softplus(-positive_score) plus F.softplus(negative_scores).sum().

Requirements

Constraints

Starter Code

import torch
import torch.nn.functional as F

def sgns_loss(center_vec: torch.Tensor, pos_vec: torch.Tensor,
              neg_vecs: torch.Tensor) -> torch.Tensor:
    """
    Returns the scalar float64 SGNS loss.
    """
    pass

Test Cases

CaseMatches
Two negativespublic
One negativepublic
Three negativespublic