๐ค AI Summary
In large-scale dual-tower retrieval, in-batch negative sampling causes over-penalization of popular items. While existing logQ correction adjusts logits to mitigate this bias, its theoretical derivation overlooks the critical fact that the positive sample always appears with probability one in the softmax denominatorโleading to incomplete correction. This work identifies this theoretical gap and proposes a refined logQ method: it explicitly models the deterministic presence of the positive sample and introduces an importance-weighting mechanism based on misclassification probability. Integrating importance sampling with Monte Carlo estimation, our approach yields a strictly unbiased gradient estimator for the softmax loss. Evaluated on multiple public and industrial benchmarks, the method consistently outperforms standard logQ, improving recall by +1.2โ2.8%. It further exhibits strong interpretability and training stability.
๐ Abstract
Two-tower neural networks are a popular architecture for the retrieval stage in recommender systems. These models are typically trained with a softmax loss over the item catalog. However, in web-scale settings, the item catalog is often prohibitively large, making full softmax infeasible. A common solution is sampled softmax, which approximates the full softmax using a small number of sampled negatives.
One practical and widely adopted approach is to use in-batch negatives, where negatives are drawn from items in the current mini-batch. However, this introduces a bias: items that appear more frequently in the batch (i.e., popular items) are penalized more heavily.
To mitigate this issue, a popular industry technique known as logQ correction adjusts the logits during training by subtracting the log-probability of an item appearing in the batch. This correction is derived by analyzing the bias in the gradient and applying importance sampling, effectively twice, using the in-batch distribution as a proposal distribution. While this approach improves model quality, it does not fully eliminate the bias.
In this work, we revisit the derivation of logQ correction and show that it overlooks a subtle but important detail: the positive item in the denominator is not Monte Carlo-sampled - it is always present with probability 1. We propose a refined correction formula that accounts for this. Notably, our loss introduces an interpretable sample weight that reflects the model's uncertainty - the probability of misclassification under the current parameters. We evaluate our method on both public and proprietary datasets, demonstrating consistent improvements over the standard logQ correction.