You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

求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:

  1. Extract context window: Grab all tokens within your chosen window size around w_t (e.g., window size 4 means 4 tokens before and after w_t, excluding w_t itself)
  2. 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 k negative tokens (w_n) that do NOT appear in w_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] and emb_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)

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):

  1. 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
  2. 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

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:

  1. 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;
      }
      
  2. 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 float arrays for embedding matrices (faster than Double for 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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.29 07:48:02