feat(m3): losses — masked LM CE + symmetric InfoNCE (T3.1)

- src/tsmm/train/losses.py:
  * lm_loss(logits, labels): shifted masked CE (answer span only); clean 0.0
    when all positions masked.
  * infonce_contrastive_loss(z1, z2): symmetric CLIP-style InfoNCE in TS-
    embedding space; positive = original vs light-augmentation view.
- tests/test_losses.py: 9 tests. 100 tests passing.

Spec clarification (small tier): InfoNCE operates in TS-embedding space
(encoder/projector output) rather than answer-embedding space; light
augmentations are positives. T3.3 strong-perturbation sensitivity is a
separate downstream check.
This commit is contained in:
张宗平
2026-06-30 02:19:02 +00:00
parent 133bbac316
commit 204d5237ba
3 changed files with 151 additions and 1 deletions
+80
View File
@@ -0,0 +1,80 @@
"""Tests for train/losses.py (T3.1)."""
import pytest
import torch
from tsmm.train.losses import lm_loss, infonce_contrastive_loss
class TestLMLoss:
def test_scalar_and_finite(self):
logits = torch.randn(2, 10, 100)
labels = torch.full((2, 10), -100, dtype=torch.long)
labels[:, 5:] = torch.randint(0, 100, (2, 5))
loss = lm_loss(logits, labels)
assert loss.dim() == 0
assert torch.isfinite(loss)
def test_all_masked_returns_zero(self):
logits = torch.randn(2, 10, 100)
labels = torch.full((2, 10), -100, dtype=torch.long)
loss = lm_loss(logits, labels)
assert float(loss) == 0.0
def test_shift_correctness(self):
# CE with shift: logits[t] predicts labels[t+1]
torch.manual_seed(0)
V = 5
logits = torch.zeros(1, 3, V)
# make position 0 strongly predict token 2
logits[0, 0, 2] = 10.0
labels = torch.full((1, 3), -100, dtype=torch.long)
labels[0, 1] = 2 # label at position 1 (shifted: predicted from logits[0])
loss = lm_loss(logits, labels)
assert float(loss) < 0.01 # near-zero when prediction is confident & correct
def test_gradient_flows(self):
logits = torch.randn(2, 8, 50, requires_grad=True)
labels = torch.full((2, 8), -100, dtype=torch.long)
labels[:, 4:] = torch.randint(0, 50, (2, 4))
loss = lm_loss(logits, labels)
loss.backward()
assert logits.grad is not None
assert logits.grad.abs().sum() > 0
class TestInfoNCE:
def test_scalar_finite(self):
z1 = torch.randn(4, 16)
z2 = torch.randn(4, 16)
loss = infonce_contrastive_loss(z1, z2)
assert loss.dim() == 0
assert torch.isfinite(loss)
def test_identical_pairs_low_loss(self):
torch.manual_seed(0)
z = torch.randn(8, 32)
z = torch.nn.functional.normalize(z, dim=-1)
loss_same = infonce_contrastive_loss(z, z, temperature=0.1)
loss_rand = infonce_contrastive_loss(z, torch.randn(8, 32), temperature=0.1)
assert float(loss_same) < float(loss_rand)
def test_symmetric(self):
z1 = torch.randn(6, 24)
z2 = torch.randn(6, 24)
a = infonce_contrastive_loss(z1, z2)
b = infonce_contrastive_loss(z2, z1)
assert torch.allclose(a, b, atol=1e-5)
def test_gradient_flows(self):
z1 = torch.randn(4, 16, requires_grad=True)
z2 = torch.randn(4, 16, requires_grad=True)
loss = infonce_contrastive_loss(z1, z2)
loss.backward()
assert z1.grad is not None and z2.grad is not None
def test_l2_normalized(self):
# loss should be invariant to per-sample scale (normalize inside)
z1 = torch.randn(4, 16)
loss_a = infonce_contrastive_loss(z1, z1)
loss_b = infonce_contrastive_loss(z1 * 5.0, z1 * 5.0)
assert torch.allclose(loss_a, loss_b, atol=1e-4)