Files
ts-as-modality/tests/test_projector.py
T
张宗平 f20c2db803 feat(m2): Projector Linear(256→896)+LayerNorm (T2.2)
- src/tsmm/model/projector.py: Projector (Linear + LayerNorm).
- tests/test_projector.py: 5 tests (shape, Qwen hidden alignment,
  LayerNorm present, gradient flow, no-NaN).
- 66 tests passing.
2026-06-30 00:17:47 +00:00

36 lines
1.1 KiB
Python

"""Tests for Projector (T2.2): TS token d → LLM hidden."""
import torch
from tsmm.model.projector import Projector
class TestProjector:
def test_output_shape(self):
proj = Projector(in_dim=256, out_dim=896)
x = torch.randn(2, 127, 256)
out = proj(x)
assert out.shape == (2, 127, 896)
def test_aligns_qwen_hidden(self):
# Qwen2.5-0.5B hidden = 896
proj = Projector(in_dim=256, out_dim=896)
out = proj(torch.randn(1, 127, 256))
assert out.shape[-1] == 896
def test_has_layernorm(self):
proj = Projector(in_dim=256, out_dim=896)
assert any(isinstance(m, torch.nn.LayerNorm) for m in proj.modules())
def test_gradient_flows(self):
proj = Projector(in_dim=256, out_dim=896)
x = torch.randn(1, 31, 256)
out = proj(x)
out.sum().backward()
assert proj.linear.weight.grad is not None
assert proj.linear.weight.grad.abs().sum() > 0
def test_no_nan(self):
proj = Projector(in_dim=256, out_dim=896)
out = proj(torch.randn(4, 128, 256))
assert not torch.isnan(out).any()