diff options
| author | Void Agent <void@jayrup.hermes> | 2026-08-02 14:20:05 +0100 |
|---|---|---|
| committer | Void Agent <void@jayrup.hermes> | 2026-08-02 14:20:05 +0100 |
| commit | 2ba0e14c3559e5786c324a89f26f159363a230b5 (patch) | |
| tree | 401683167a9541e066c3a30c83037bf7cc89e83a | |
| parent | 42660815eef000bcb662287a97cbb3012b6c90b8 (diff) | |
Document n_pairs denominator (exact count) and ctrl_random gradient-mass caveat (comment-only)
| -rw-r--r-- | src/jlens_v3.py | 3 | ||||
| -rw-r--r-- | src/loss_reweight.py | 5 |
2 files changed, 8 insertions, 0 deletions
diff --git a/src/jlens_v3.py b/src/jlens_v3.py index b5f59ed..69299dc 100644 --- a/src/jlens_v3.py +++ b/src/jlens_v3.py @@ -115,6 +115,9 @@ def compute_faithful_jlens(model, layer_idx, batches, device, chunk=32): torch.cuda.empty_cache() n_pairs += B * T * (T + 1) // 2 # valid (t, t' >= t) pairs + # NOTE: count = T(T+1)/2 per sequence (all source positions t and all + # futures t' >= t). Pairs with t' < t contribute zero gradient by + # causality, so summing over all t' and dividing by this count is exact. del logits, loss, h_l, h_final, h_final_sum torch.cuda.empty_cache() 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) |
