求Java版Word2Vec Skip-Gram算法,需实现支持n-gram的负采样模型
Hey there! Let's start with the core Skip-Gram with Negative Sampling (SGNS) algorithm you're looking for, then walk through how to tweak it to support n-gram contexts—since you're already set up with Java and have your corpus processed, this should fit right in.
Core Skip-Gram with Negative Sampling Algorithm
First, let's break down the vanilla SGNS logic, including the sigmoid component you need:
1. Preprocessing (You’ve already handled this, but aligning for clarity)
- Convert your corpus into a sequence of tokens
- Build a vocabulary of unique tokens, assigning each a unique integer ID
- Calculate token frequencies to power negative sampling (we’ll use this to sample "irrelevant" tokens later)
2. Training Loop (Per Target Token)
For every target token w_t in your corpus:
- Extract context window: Grab all tokens within your chosen window size around
w_t(e.g., window size 4 means 4 tokens before and afterw_t, excludingw_titself) - Process each context element:
a. Positive Pair Calculation:- Compute the dot product of the target token’s embedding (
emb_target[w_t]) and the context token’s embedding (emb_context[w_c]) - Run the result through the sigmoid function to get a probability of the pair being a valid context-target pair:
double sigmoid(double x) { if (x > 10) return 1.0; // Avoid overflow if (x < -10) return 0.0; return 1.0 / (1.0 + Math.exp(-x)); } - The loss for this positive pair is
-log(sigmoid(dot_product))—we want to maximize this probability
b. Negative Sampling: - Sample
knegative tokens (w_n) that do NOT appear inw_t’s context window. Weight samples by token frequency (use a cumulative frequency array for fast sampling in Java)
c. Negative Pair Calculation: - For each negative token, compute the dot product of
emb_target[w_t]andemb_context[w_n], then apply sigmoid to the negative result:sigmoid(-dot_product) - The loss for each negative pair is
-log(sigmoid(-dot_product))—we want to minimize the probability of these invalid pairs
d. Update Embeddings: - Use stochastic gradient descent (SGD) to adjust both target and context embeddings to reduce the total loss (sum of positive and negative pair losses)
- Compute the dot product of the target token’s embedding (
Adapting to N-Gram Contexts
To replace individual context tokens with n-grams, you have two solid approaches depending on your goals:
Approach 1: Treat N-Grams as Unique Vocabulary Entries
This works if you want to learn distinct embeddings for each n-gram (e.g., bigrams like "New York" get their own vector):
- Preprocessing Update:
- Generate all valid n-grams from your corpus (slide an n-length window over the token sequence)
- Extend your vocabulary to include these n-grams, assigning each a unique ID just like individual tokens
- Update your frequency counts to track n-gram occurrences for negative sampling
- Training Loop Modification:
- When extracting the context window around
w_t, generate all n-grams within that window instead of individual tokens - Use these n-grams as your context elements (
g_c) in the training loop—everything else (positive/negative loss, embeddings update) stays the same as the core SGNS algorithm
- When extracting the context window around
Approach 2: Combine Token Embeddings for N-Grams
If you want to avoid expanding your vocabulary (great for larger n or dynamic n-gram sizes), represent each n-gram as a combination of its constituent token embeddings:
- Preprocessing Update:
- No need to add n-grams to your vocabulary; keep using your existing token vocabulary
- For any n-gram, compute its context vector by averaging (or summing) the embeddings of each token in the n-gram:
float[] getNGramEmbedding(List<Integer> nGramTokenIds, float[][] embContext) { float[] nGramEmb = new float[embContext[0].length]; for (int tokenId : nGramTokenIds) { for (int i = 0; i < nGramEmb.length; i++) { nGramEmb[i] += embContext[tokenId][i]; } } // Normalize by n to get average for (int i = 0; i < nGramEmb.length; i++) { nGramEmb[i] /= nGramTokenIds.size(); } return nGramEmb; }
- Training Loop Modification:
- When processing a context n-gram, use this combined vector instead of a single token embedding in the dot product step
- Negative sampling can still use individual tokens (or you can sample n-grams if you prefer)—the loss calculation and embedding updates remain consistent
Java-Specific Tips
Since you’re building this in Java, a few practical pointers:
- Use 2D
floatarrays for embedding matrices (faster thanDoublefor large vocabularies) - Precompute a cumulative frequency array for negative sampling: this lets you use binary search to sample tokens/n-grams in O(log V) time instead of O(V)
- Batch your training updates if possible to reduce overhead, especially with large corpora
内容的提问来源于stack exchange,提问作者Liran

