summaryrefslogtreecommitdiff
path: root/src/loss_reweight.py
diff options
context:
space:
mode:
authorVoid Agent <void@jayrup.hermes>2026-08-02 14:20:05 +0100
committerVoid Agent <void@jayrup.hermes>2026-08-02 14:20:05 +0100
commit2ba0e14c3559e5786c324a89f26f159363a230b5 (patch)
tree401683167a9541e066c3a30c83037bf7cc89e83a /src/loss_reweight.py
parent42660815eef000bcb662287a97cbb3012b6c90b8 (diff)
Document n_pairs denominator (exact count) and ctrl_random gradient-mass caveat (comment-only)
Diffstat (limited to 'src/loss_reweight.py')
-rw-r--r--src/loss_reweight.py5
1 files changed, 5 insertions, 0 deletions
diff --git a/src/loss_reweight.py b/src/loss_reweight.py
index 1bf3124..f01d88a 100644
--- a/src/loss_reweight.py
+++ b/src/loss_reweight.py
@@ -53,6 +53,11 @@ def _weighted_loss(logits, y, mode, q_id, batch_k, V, device):
if mode == 'q':
w[y.view(-1) == q_id] = WEIGHT
elif mode == 'ctrl_random':
+ # Same number of upweighted positions as the q-mode model, but on
+ # random non-q targets. NOTE: gradient-magnitude distribution differs
+ # from q-mode (random positions spread across the batch vs rare 'q'
+ # positions); the design controls for "any 2x reweighting changes the
+ # model", not for exact gradient-mass matching.
g = torch.Generator().manual_seed(1000 + batch_k) # CPU generator (randperm)
n_q = int((y == q_id).sum().item())
flat = torch.arange(y.numel(), device=device)