summaryrefslogtreecommitdiff
path: root/src/loss_reweight.py
diff options
context:
space:
mode:
authorVoid Agent <void@jayrup.hermes>2026-08-02 13:32:16 +0100
committerVoid Agent <void@jayrup.hermes>2026-08-02 13:32:16 +0100
commit53de9ecd13fbf7e098fea88f73d1a9403a40d83b (patch)
tree7d6e1289f19a2370c9fae2a339ad1d263f074f87 /src/loss_reweight.py
parent30e81b39e1e280c8fe51e13349ebec626a5353aa (diff)
loss_reweight: pass Y to forward for full logits (Karpathy nanoGPT returns last-position logits when targets=None); regression test
Diffstat (limited to 'src/loss_reweight.py')
-rw-r--r--src/loss_reweight.py6
1 files changed, 3 insertions, 3 deletions
diff --git a/src/loss_reweight.py b/src/loss_reweight.py
index bcfdf07..1bf3124 100644
--- a/src/loss_reweight.py
+++ b/src/loss_reweight.py
@@ -108,8 +108,8 @@ def train(mode, seed, max_iters, batch_size):
for _ in range(50):
X, Y = get_batch('val')
with torch.no_grad():
- logits = model(X)[0]
- lv.append(F.cross_entropy(logits.view(-1, V), Y.view(-1)).item())
+ _, loss = model(X, Y)
+ lv.append(loss.item())
v = np.mean(lv)
model.train()
if v < best_val:
@@ -119,7 +119,7 @@ def train(mode, seed, max_iters, batch_size):
if it % 1000 == 0:
print(f" iter {it}: val={v:.4f}")
X, Y = get_batch('train')
- logits = model(X)[0]
+ logits = model(X, Y)[0] # pass Y: targets=None would give last-position logits only
loss = _weighted_loss(logits, Y, mode, q_id, it, V, device)
loss.backward()
opt.step()